继承sklearn BaseEstimator的自定义类属性名疑问及报错解析
自定义scikit-learn Estimator:属性名与__init__参数名的问题解答
基础疑问解答
编写继承自scikit-learn BaseEstimator的自定义类时,对于需要被scikit-learn核心API识别的参数,必须保证__init__方法的参数名与实例属性名完全一致。这是因为BaseEstimator依赖这种一一对应关系来实现参数获取(get_params)、参数设置(set_params)、模型序列化等核心功能。
具体问题分析
报错原因
你的代码中,__init__方法定义了参数row_wise,但你将其赋值给了实例属性self.row,而非同名的self.row_wise。当执行print(null)或访问实例参数时,scikit-learn会调用get_params方法,该方法会遍历__init__的所有参数名,并尝试通过getattr(self, 参数名)获取对应属性值。由于不存在self.row_wise这个属性,直接触发AttributeError。
另外,你的代码还遗漏了将variables参数赋值给同名实例属性self.variables,这同样会在后续调用get_params时引发问题。
是否可以使用不同的属性名存储参数值?
可以,但必须保留与__init__参数名完全一致的实例属性,额外的别名属性可以按需添加。也就是说,你可以同时保留self.row_wise = row_wise和self.row = row_wise,既满足scikit-learn的机制要求,又能使用自定义的别名属性。
修改后的代码示例
from sklearn.base import BaseEstimator, TransformerMixin from typing import Union, List class BaseNull(BaseEstimator, TransformerMixin): def __init__(self, variables: Union[str, List[str]], row_wise: bool = False, na_kwds: Union[str, List, tuple] = None): # 必须保留与__init__参数同名的属性,满足scikit-learn要求 self.variables = variables self.row_wise = row_wise self.na_kwds = na_kwds # 自定义别名属性(可选) self.row = row_wise # 其他内部属性 self.null_index = None self.null_columns = None self.row_count = None self.column_count = None
关键总结
BaseEstimator的get_params/set_params方法严格依赖__init__参数名与实例属性名的对应关系- 所有
__init__中定义的公共参数,都必须赋值给同名的实例属性 - 自定义别名属性仅作为补充,不能替代同名属性的定义
内容的提问来源于stack exchange,提问作者Hewarathna Ashen Iranga
相关产品推荐
相关产品推荐

