4.1 自定义估计器与转换器


文档摘要

4.1 自定义估计器与转换器 本节摘要:自定义估计器与转换器是扩展 Scikit-learn 的核心手段。估计器是实现了 fit 和 predict 的对象,负责从数据中学习;转换器是实现了 fit 和 transform 的对象,负责把数据加工成新形态。当你遇到内置组件覆盖不了的业务逻辑时,就继承 BaseEstimator 和 TransformerMixin(或 RegressorMixin、ClassifierMixin),补上这几个方法,你的类就能立刻获得 getparams、setparams、fittransform 等能力,并顺畅插进管道和交叉验证。 本节地图 阅读完本节,你应当能够: 区分估计器和转换器在方法签名上的本质差异。

4.1 自定义估计器与转换器

本节摘要:自定义估计器与转换器是扩展 Scikit-learn 的核心手段。估计器是实现了 fit 和 predict 的对象,负责从数据中学习;转换器是实现了 fit 和 transform 的对象,负责把数据加工成新形态。当你遇到内置组件覆盖不了的业务逻辑时,就继承 BaseEstimator 和 TransformerMixin(或 RegressorMixin、ClassifierMixin),补上这几个方法,你的类就能立刻获得 get_params、set_params、fit_transform 等能力,并顺畅插进管道和交叉验证。

本节地图

阅读完本节,你应当能够:

  1. 区分估计器和转换器在方法签名上的本质差异。
  2. 独立实现一个可用的自定义转换器,并解释 fit 为什么要返回 self。
  3. 独立实现一个自定义估计器,并正确处理未拟合时的报错。
  4. 说清 get_params 和 set_params 为什么是管道与网格搜索能工作的前提。
  5. 用 check_estimator 校验自定义组件的接口合规性。

一、先写一个能跑的,再谈为什么

别急着背概念。我们直接写一个最傻的转换器——把输入数据整体开个平方根,然后看它能不能跑起来。

import numpy as np from sklearn.base import BaseEstimator, TransformerMixin class SquareRootScaler(BaseEstimator, TransformerMixin): def __init__(self): pass def fit(self, X, y=None): return self def transform(self, X): return np.sqrt(X)

就这么点代码。实例化后调 fit_transform,你能直接拿到结果——TransformerMixin 已经帮你把 fit_transform 拼好了,它内部就是先 fittransform。这第一段代码想让你记住一件事:自定义组件不是黑魔法,它只是"继承两个基类 + 填两个方法"的体力活。真正的难点在后面——怎么让这个体力活做对、做得不埋雷。

那为什么 fit 里什么都不干,还要写一行 return self?因为 Scikit-learn 全生态都默认"先 fit 再 transform"这个动作序列。管道会机械地对每个步骤先调 fit,你的 fit 不返回 self,管道下一步就拿不到这个对象,直接报错。这个约定比你想的更重要,我们第三节再展开。

二、估计器和转换器到底差在哪

把上面那段代码看懂之后,我们把抽象的定义补上。Scikit-learn 里几乎所有对象都归到两类里:

  • 估计器:实现了 fit(X, y)predict(X),核心职责是从数据里"学"出参数。分类器、回归器、聚类器都算估计器。
  • 转换器:实现了 fit(X, y)transform(X),核心职责是把数据从一种形态变成另一种形态。标准化、独热编码、主成分分析都算转换器。

注意一个容易绕晕的点:转换器也有 fit,但它的 fit 不是"训练模型",而是"学习转换所需的统计量"。比如标准化器的 fit 就是算出均值和标准差存起来,transform 时才真正用它们去减均值、除标准差。

两者的方法对照,一张表说清楚:

方法 估计器 转换器
fit 学习模型参数 学习转换统计量(均值、分位数等)
transform 通常没有 用统计量加工数据
predict 输出预测值 通常没有
fit_transform 通常没有 由 Mixin 提供,先 fit 再 transform
score 评估性能,Mixin 提供默认实现 通常没有

为什么要做这样的划分?因为 Scikit-learn 的设计哲学是"所有东西长得一样"。你只要知道一个对象是估计器还是转换器,就能猜出它能做什么、怎么插进管道。这个一致性让交叉验证、网格搜索、管道这些高层工具可以不分青红皂白地操作任何组件。

三、三个方法,各有各的坑

先看一条贯穿全程的数据流——不管你的组件多复杂,运行时都逃不出"先 fit 再 transform 或 predict"这个动作序列。训练集喂给 fit,学出统计量或参数;测试集走 transform 加工、或走 predict 预测。这条链是管道能自动串起一切的底层契约。

3.1 fit 必须返回 self

我们已经见过 return self。它背后是一条铁律:fit 修改对象自身的状态,并把这个对象本身返回。为什么要这样?管道和网格搜索会写类似这样的代码——step = step.fit(X, y),然后把返回值接着往下传。如果你的 fit 返回了别的,整个链就断了。

3.2 transform 不许改输入

transform 的第一行最好就复制一份数据,别在传入的数组上原地改。调用方很可能还要用原始数据做别的,你悄悄改了,就是个难查的副作用。

def transform(self, X): X = np.copy(X) # 先复制,避免污染调用方的数据 X[:, self.target_col] = self.encoder(X[:, self.target_col]) return X

3.3 predict 要拦住"没 fit 就预测"

估计器有个经典错误:对象还没 fit,就有人调 predict。你不拦,它要么崩得莫名其妙,要么静默返回垃圾。正确做法是抛 NotFittedError,提示调用方先 fit。

from sklearn.exceptions import NotFittedError class MeanRegressor(BaseEstimator, RegressorMixin): def fit(self, X, y): self.mean_ = np.mean(y) return self def predict(self, X): if not hasattr(self, "mean_"): raise NotFittedError("先调 fit 再 predict") return np.full(X.shape[0], self.mean_)

注意一个命名细节:fit 里学习出来的属性,我习惯在名字末尾加下划线,比如 mean_coef_。这是 Scikit-learn 的约定,用于区分"用户传的超参数"和"模型学出来的参数"。超参数不加下划线(存在 __init__ 里),学出来的加下划线。这个习惯不只是好看,它还跟 get_params 的行为绑定,见下一节。

图:自定义估计器结构

图:自定义估计器结构

3.4 fit 不是摆设:一个真正会"学"的转换器

前面两个例子的 fit 都是空转,容易让人误以为 fit 只是个形式。其实转换器的 fit 常常要真干活。写一个"分位数裁剪器"——把超出上下分位数的极端值夹回边界,这个边界就得靠 fit 从训练数据里算出来。

class QuantileClipper(BaseEstimator, TransformerMixin): def __init__(self, lower=0.01, upper=0.99): self.lower = lower self.upper = upper def fit(self, X, y=None): X = np.asarray(X) self.lower_ = np.percentile(X, self.lower * 100, axis=0) self.upper_ = np.percentile(X, self.upper * 100, axis=0) return self def transform(self, X): X = np.asarray(X) return np.clip(X, self.lower_, self.upper_)

看这几个细节:lowerupper 是超参数,原样存在 __init__ 里,没有下划线;lower_upper_ 是 fit 学出来的分位数,带下划线。fit 里没有碰测试数据,只从训练集算边界;transform 拿这个边界去夹测试集。整个对象放进管道,行为就和内置的 RobustScaler 一样规矩。这才是自定义转换器的标准姿势——fit 负责从训练集学统计量,transform 负责把统计量一致地施加到任何数据上

四、get_params 和 set_params:为什么它们是管道的命门

你可能没直接调过 get_params,但网格搜索每天都在调它。它的规矩很死:get_params 要返回 __init__ 里每一个超参数,且键名和参数名一模一样。你 __init__ 里写 def __init__(self, degree=2),get_params 就得返回 {"degree": 2}

这套规矩存在的意义,是让网格搜索能这样工作:它先 get_params 拿到当前超参,改一个值,再用 set_params 塞回去,重新 fit。如果你在 __init__ 里把参数改了个名、或者把 degree 偷偷转成了 self.deg,get_params 和 set_params 就对不上账,网格搜索立刻抓瞎。所以老规矩:超参数在 __init__ 里原样存成同名属性,一个字都别改

BaseEstimator 已经帮你把这两个方法实现了,前提是你遵守"同名存储"这条约定。你几乎不用手写它们,但得知道它们的脾气,否则出问题时一头雾水。

⚠️ 常见坑:在 __init__ 里做重活,比如 self.scaler_ = StandardScaler().fit(X)。构造函数应该只做"参数存起来"这一件事,任何碰数据的计算都放到 fit 里。否则克隆估计器时,构造函数会被反复调用,轻则慢,重则状态错乱。

💡 关键直觉:把超参数想象成"配方",把 fit 学出来的东西想象成"做好的菜"。配方必须能原样誊抄(get_params/set_params),菜只能在做菜环节(fit)产生。两者分家,克隆、调参、部署才不会串味。

五、让自定义组件过一遍"体检"

写完组件,别急着上线,先过一遍 Scikit-learn 自带的体检工具 check_estimator。它会跑一堆针对接口规范的测试,检查你的类在各种刁钻输入下会不会崩、会不会漏方法。

from sklearn.utils.estimator_checks import check_estimator check_estimator(SquareRootScaler)

它通过,只说明"接口合规",不说明"算法正确"。接口合规是及格线,算法正确还得靠你自己的单元测试和真实数据验证。这两件事别混为一谈——体检能抓"没写 predict 就敢说自己是估计器"这种低级错误,抓不住"你算出来的均值根本不对"这种逻辑错误。

六、什么时候值得自己造

不是所有需求都值得手写。我的判断标准很朴素:先查内置组件,再查能不能用管道组合出来,最后才考虑自己写。三者都覆盖不了,才动手造。

典型的"该自己造"场景有这些:特定领域的清洗逻辑(比如基因序列、时间窗口的滑动统计)、把已有的老代码包成标准组件、实现一个研究里刚提出的算法。典型的"不该自己造"场景:标准化、独热编码这种内置现成的,别重复造轮子,徒增维护成本。判断标准就一句:先查内置、再查组合、最后才自造。

自己造的额外好处是"可组合性"。你写的转换器能塞进管道,能参与网格搜索,能跟内置组件平起平坐。这个好处不是白来的,它是你前面遵守 fit 返回 self、超参同名存储、transform 不污染输入这三条规矩换来的回报。

from sklearn.pipeline import Pipeline from sklearn.model_selection import GridSearchCV from sklearn.preprocessing import PolynomialFeatures from sklearn.linear_model import LogisticRegression pipe = Pipeline([ ("clip", QuantileClipper()), ("poly", PolynomialFeatures(degree=2)), ("model", LogisticRegression()), ]) grid = GridSearchCV(pipe, {"clip__lower": [0.0, 0.01], "model__C": [0.1, 1.0]}, cv=5) grid.fit(X_train, y_train)

注意 clip__lower 这个双下划线写法——它穿过管道,直接定位到 QuantileClipper 的 lower 超参数。你之所以能被网格搜索这样精确地调参,正是因为前面守住了"超参同名存储"这条规矩。自造组件和内置组件在这里没有区别,都吃同一套调参语法。

常见疑问

为什么我的转换器要同时继承两个基类?
BaseEstimator 给参数管理能力,TransformerMixin 给 fit_transform 的默认实现。缺了前者,网格搜索调不了参;缺了后者,你得多写一个 fit_transform。两个都要,才是一个"满血"的转换器。

自定义估计器和自定义转换器哪个更难?
估计器更难一点。转换器只需管好 fit 和 transform 的数据流,估计器还要管 predict 的输出形状、score 的评估逻辑,以及未拟合时的报错。建议先从转换器入手,摸熟了再碰估计器。

clone 一个自定义组件会丢状态吗?
不会,前提是你守了"超参同名存储"这条规矩。Scikit-learn 的 clone 机制会读 get_params 拿到所有超参,用它们重新构造一个全新的、未 fit 的对象——所以构造阶段绝不能碰数据、不能有副作用,否则克隆出来的对象就脏了。

fit 里的 y 参数有什么用?转换器要它干嘛?
很多转换器不需要目标变量,fit 的 y 参数就设成默认 None、直接忽略。保留这个参数是为了接口统一——管道在调 fit 时会机械地传 X 和 y,你的 fit 不接受 y 就会崩。所以哪怕用不上,签名里也要留着。

怎么判断自己的实现对不对?
除了 check_estimator 验接口,还要做两件事:用小数据手算一遍,确认 fit 学出的统计量没错;再拿它和内置的近似组件对比结果。比如你自己写的标准化器,输出应该和 StandardScaler 完全一致,对不上就说明有 bug。

要点串联

  • 估计器与转换器的分野:前者 fit 加 predict,学模型;后者 fit 加 transform,加工数据。
  • fit 必须返回 self:管道和网格搜索都依赖这个链式调用约定。
  • transform 先复制再改:避免污染调用方的数据,杜绝难查的副作用。
  • 超参数同名存储__init__ 里原样存,get_params 和 set_params 才能对上账。
  • 学出的属性加下划线:区分"配方"和"做好的菜",跟 Scikit-learn 约定一致。
  • predict 未拟合要报错:抛 NotFittedError,而不是静默返回垃圾。
  • check_estimator 只验接口:算法正确性还得靠自己的测试和真实数据。

下一节我们把镜头从"单个零件"拉到"一组零件的协作"——集成方法。你会看到,刚才学会的自定义估计器,正是集成方法里最灵活的基学习器来源。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U