如何让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
相关产品推荐
相关产品推荐

