使用Scikit-learn包的TransformerMixin
TransformerMixin 是 Scikit-learn 提供的一个混入类 (Mixin class),它为你创建自定义的数据转换器提供了极大的便利。
它的核心作用是自动为你实现 fit_transform 方法,让你无需重复编写样板代码。
TransformerMixin 是什么?
在 Scikit-learn 中,转换器(Transformer)是用于数据预处理或特征工程的工具,它们通常都有 fit()、transform() 和 fit_transform() 这三个方法。
TransformerMixin 就是一个提供 fit_transform 方法实现的混入类。它的实现逻辑非常直接:
- 先调用你定义的
fit()方法。 - 再调用你定义的
transform()方法。
这意味着,只要你写的类继承了 TransformerMixin,并正确实现了 fit 和 transform 方法,就能免费获得一个标准的 fit_transform 方法。
为什么需要它?
它主要解决了两个问题:
- 遵循 Scikit-learn 的 API 标准:继承
TransformerMixin能确保你创建的自定义转换器与 Scikit-learn 生态系统(如Pipeline)完美兼容。 - 代码更简洁:你不需要在每个自定义转换器里重复编写
fit_transform方法。
如何使用:一个完整的示例
通常,一个自定义转换器会同时继承 TransformerMixin 和 BaseEstimator。BaseEstimator 提供了 get_params() 和 set_params() 方法,这对于在 Pipeline 或 GridSearchCV 中使用至关重要。
下面是一个简单的自定义转换器示例,它会在训练时记录一个参数,并在转换时用该参数填充一个数组:
import numpy as np
from sklearn.base import BaseEstimator, TransformerMixin
# 1. 定义类,同时继承 TransformerMixin 和 BaseEstimator
class MyTransformer(TransformerMixin, BaseEstimator):
def __init__(self, param=1):
self.param = param
# 2. 实现 fit 方法,必须返回 self
def fit(self, X, y=None):
# 这里可以学习数据中的参数,此示例只是简单返回自身
return self
# 3. 实现 transform 方法,返回转换后的数据
def transform(self, X):
# 此示例简单地返回一个填充了 self.param 的数组
return np.full(shape=len(X), fill_value=self.param)
# 使用自定义转换器
transformer = MyTransformer()
X = [[1, 2], [2, 3], [3, 4]]
# 可以直接使用 fit_transform 方法
result = transformer.fit_transform(X)
print(result) # 输出: array([1, 1, 1])
核心方法与高级功能
除了自动实现 fit_transform,TransformerMixin 还提供了其他重要功能:
fit_transform(X, y=None, **fit_params):这是最主要的方法,功能已在上文介绍。set_output(*, transform=None):这是一个较新加入的功能,允许你控制transform和fit_transform方法的输出类型。你可以将其设置为返回 Pandas DataFrame 或 Polars DataFrame,而不仅仅是 NumPy 数组。
注意事项
- 继承顺序:虽然示例中写作
class MyTransformer(TransformerMixin, BaseEstimator):,但在 Python 多重继承中,Mixin 类通常放在前面。 - 何时需要覆盖
fit_transform:TransformerMixin提供的默认实现是先fit再transform。如果你的转换器有更高效的方法可以同时进行拟合和转换(例如StandardScaler),你可以选择在自己的类中重写(覆盖)这个方法。
总结
简而言之,TransformerMixin 是构建 Scikit-learn 兼容的自定义转换器不可或缺的助手。它通过提供标准的 fit_transform 实现,让你能专注于核心的 fit 和 transform 逻辑,从而写出更干净、更标准的代码。
- 点赞
- 收藏
- 关注作者
评论(0)