构建Scikit-learn兼容估计器遇AttributeError: 'dict'无'requires_fit'属性
解决Scikit-learn自定义估计器的AttributeError问题
错误原因分析
出现AttributeError: 'dict' object has no attribute 'requires_fit'的核心原因是:
- 在scikit-learn 1.6.1版本中,
__sklearn_tags__必须定义为类方法,而非实例方法。直接返回字典会导致内部代码尝试将字典当作Tag对象访问属性,从而触发错误。 - 代码还存在两处次要错误:
predict方法中误写self.intercept,实际应为self.intercept_(fit方法中定义的是带下划线的属性)。_obtain_beta方法中self.coef_的维度错误:误用了样本数X.shape[0],实际应使用特征数X.shape[1],否则矩阵相乘会出现维度不匹配。
修正后的完整代码
from sklearn.utils.validation import check_is_fitted, check_X_y import numpy as np from sklearn.base import BaseEstimator, RegressorMixin class BaseModel(BaseEstimator, RegressorMixin): """ Base class for penalized regression models using cp. """ def __init__(self, param: float = 0.5): self.param = param def _obtain_beta(self, X, y): self.intercept_ = 0 # 修正:用特征数X.shape[1]而非样本数X.shape[0] self.coef_ = self.param * np.ones(X.shape[1]) def fit(self, X: np.ndarray, y: np.ndarray): self.feature_names_in_ = None if hasattr(X, "columns"): self.feature_names_in_ = np.asarray(X.columns, dtype=object) X, y = check_X_y(X, y, accept_sparse=False, y_numeric=True, ensure_min_samples=2) self.n_features_in_ = X.shape[1] self._obtain_beta(X, y) self.is_fitted_ = True return self def predict(self, X: np.ndarray) -> np.ndarray: check_is_fitted(self, ["coef_", "intercept_", "is_fitted_"]) # 修正:用self.intercept_而非self.intercept predictions = np.dot(X, self.coef_) + self.intercept_ return predictions @classmethod # 关键:将__sklearn_tags__定义为类方法 def __sklearn_tags__(cls): # 先获取父类的标签,再合并自定义标签 tags = super().__sklearn_tags__() tags.update({ "allow_nan": False, "requires_y": True, "requires_fit": True, }) return tags # USAGE EXAMPLE from sklearn.datasets import make_regression X, y, beta = make_regression(n_samples=200, n_features=200, n_informative=25, bias=10, noise=5, random_state=42, coef=True) model = BaseModel() model.fit(X, y) # 验证预测功能 preds = model.predict(X) print(preds[:5])
关键修正点说明
@classmethod装饰器:确保__sklearn_tags__是类方法,scikit-learn内部会正确调用并处理返回的标签字典。- 继承父类标签:通过
super().__sklearn_tags__()获取BaseEstimator和RegressorMixin的默认标签,再用update添加自定义标签,避免覆盖必要的默认标签。 - 维度与属性名修正:解决矩阵相乘和属性不存在的潜在错误。
内容的提问来源于stack exchange,提问作者Álvaro Méndez Civieta
相关产品推荐
相关产品推荐

