You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何让scikit-learn自定义estimator在训练和预测阶段执行不同处理逻辑

实现思路

scikit-learn的Transformer执行逻辑天然可以区分训练和预测场景:训练阶段会优先调用fit_transform方法,预测阶段只会调用transform方法,我们可以利用这个规则重写对应方法,不需要额外传参就能自动适配两种场景的预处理逻辑,同时完全兼容Pipeline和GridSearchCV。

修改后的代码实现

from sklearn import preprocessing
from sklearn.base import BaseEstimator, TransformerMixin
from Levenshtein import distance # 用于编辑距离计算,可替换为自己的匹配逻辑

class TextPreprocessor(BaseEstimator, TransformerMixin):
    def __init__(self):
        self.train_vocab = set() # 存储训练集所有token
        # 注意:LabelEncoder是处理标签的工具,不建议放在特征预处理Transformer中
        # 这里保留仅为对齐原代码,建议单独拆分标签处理逻辑
        self.le = preprocessing.LabelEncoder()

    @staticmethod
    def do_basic_preprocessing(text, is_training, train_vocab=None):
        # 通用预处理步骤
        text = preprocessing_1(text)
        text = preprocessing_2(text)
        
        # 仅预测阶段执行拼写纠错
        if not is_training:
            corrected_tokens = []
            for token in text.split():
                # 找训练词汇表里编辑距离最小的token
                min_dist = float('inf')
                matched_token = token
                for vocab_token in train_vocab:
                    dist = distance(token, vocab_token)
                    if dist < min_dist:
                        min_dist = dist
                        matched_token = vocab_token
                corrected_tokens.append(matched_token)
            text = " ".join(corrected_tokens)
        
        text = preprocessing_n(text)
        return text

    def fit(self, X, y=None):
        # 存储训练集所有token,用于后续预测阶段纠错
        for text in X:
            tokens = text.split() # 可替换为自己的分词逻辑
            self.train_vocab.update(tokens)
        if y is not None:
            self.le.fit(y)
        return self

    def transform(self, X, is_training=False):
        X_processed = [
            self.do_basic_preprocessing(text, is_training=is_training, train_vocab=self.train_vocab)
            for text in X
        ]
        # 注意:Transformer的transform方法仅允许返回处理后的特征X,不能返回y,否则会导致Pipeline报错
        # 如果需要处理y,建议单独封装标签转换器,不要和特征预处理逻辑耦合
        return X_processed

    def fit_transform(self, X, y=None):
        # 训练阶段调用,自动走is_training=True的预处理逻辑
        return self.fit(X, y).transform(X, is_training=True)

注意事项

  • 训练阶段调用fit_transform时自动执行无纠错的预处理逻辑,预测阶段直接调用transform()默认走带纠错的逻辑,不需要手动传参,放到Pipeline、GridSearchCV中运行也会自动保持该行为。
  • 原代码中把标签编码逻辑放在特征预处理类里不符合scikit-learn的设计规范,会导致Pipeline运行异常,建议单独编写标签处理类,在模型训练前单独对y做处理。
  • 拼写纠错的匹配逻辑可以根据业务需求优化,比如添加编辑距离阈值,超过阈值的token直接保留原内容,避免误修正。

内容的提问来源于stack exchange,提问作者hafiz031

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.29 00:45:03