使用Scikit-learn包的RegressorMixin
【摘要】 RegressorMixin 是 Scikit-learn 中的一个混入类(Mixin class),它为所有回归估计器(regression estimators)提供标准功能。你可以把它看作一个“功能包”,将它和BaseEstimator一起继承,就能快速创建一个符合 Scikit-learn 标准的自定义回归模型。它主要提供以下核心功能:自动设置估计器类型:将_estimator_ty...
RegressorMixin 是 Scikit-learn 中的一个混入类(Mixin class),它为所有回归估计器(regression estimators)提供标准功能。你可以把它看作一个“功能包”,将它和BaseEstimator一起继承,就能快速创建一个符合 Scikit-learn 标准的自定义回归模型。
它主要提供以下核心功能:
- 自动设置估计器类型:将
_estimator_type属性设为"regressor",让 Scikit-learn 能识别出这是一个回归器。 - 提供默认的
score方法:实现了score()方法,默认使用决定系数 R²来评估模型性能。 - 强制要求目标值
y:通过requires_y标签确保fit()方法必须传入目标值y。
如何使用:创建一个自定义回归器
使用RegressorMixin的标准做法是,创建一个同时继承RegressorMixin和BaseEstimator的类。
注意:为了确保正确的方法解析顺序(MRO),RegressorMixin应该放在BaseEstimator的前面。
以下是一个官方文档中的完整示例:
import numpy as np
from sklearn.base import BaseEstimator, RegressorMixin
# 1. 定义自定义回归器类,继承顺序:RegressorMixin 在前
class MyEstimator(RegressorMixin, BaseEstimator):
def __init__(self, *, param=1):
self.param = param
# 2. 必须实现 fit 方法
def fit(self, X, y=None):
# 在这里实现模型训练逻辑
self.is_fitted_ = True # 惯例:标记模型已拟合
return self # 必须返回 self
# 3. 必须实现 predict 方法
def predict(self, X):
# 在这里实现预测逻辑
return np.full(shape=X.shape[0], fill_value=self.param)
# 使用自定义模型
estimator = MyEstimator(param=0)
X = np.array([[1, 2], [2, 3], [3, 4]])
y = np.array([-1, 0, 1])
# 拟合、预测、评分
estimator.fit(X, y)
predictions = estimator.predict(X)
score = estimator.score(X, y) # 直接使用 RegressorMixin 提供的 score 方法
print(predictions) # 输出: [0 0 0]
print(score) # 输出: 0.0
深入理解 score 方法
RegressorMixin提供的score方法,其返回值是决定系数 R²。
- 计算公式:( R^2 = 1 - \frac{u}{v} )。
u是残差平方和:((y_true - y_pred) ** 2).sum()。v是总平方和:((y_true - y_true.mean()) ** 2).sum()。
- 分数解读:最高分为
1.0,分数可以为负(表示模型预测效果极差)。一个总是预测y均值的“常数模型”,其 R² 为0.0。
总结
简单来说,创建自定义回归模型时,遵循这个模板即可:
- 导入
BaseEstimator和RegressorMixin。 - 定义新类,继承顺序为
(RegressorMixin, BaseEstimator)。 - 实现
__init__(用于设置参数)、fit(训练逻辑)和predict(预测逻辑)方法。 - 这样,你的自定义模型就自动拥有了
get_params()、set_params()(来自BaseEstimator)和score()(来自RegressorMixin)等标准方法,可以无缝融入 Scikit-learn 的生态系统(如网格搜索GridSearchCV)。
【版权声明】本文为华为云社区用户原创内容,未经允许不得转载,如需转载请自行联系原作者进行授权。如果您发现本社区中有涉嫌抄袭的内容,欢迎发送邮件进行举报,并提供相关证据,一经查实,本社区将立刻删除涉嫌侵权内容,举报邮箱:
cloudbbs@huaweicloud.com
- 点赞
- 收藏
- 关注作者
评论(0)