如何固定scikit-learn评估器参数且兼容Pipeline、GridSearchCV等工具
问题根因
scikit-learn的get_params()方法依赖解析评估器__init__方法的显式参数签名收集参数,你之前自定义子类时仅使用**kwargs作为入参,没有显式声明父类的所有参数,导致get_params()无法识别参数,返回空字典。
方案1:符合scikit-learn规范的自定义子类实现(优先推荐)
该方案完全符合scikit-learn评估器开发规范,可强制锁定目标参数不被覆盖,完美适配Pipeline、GridSearchCV等所有工具。
实现代码
from sklearn.ensemble import RandomForestClassifier class FiveTreesClassifier(RandomForestClassifier): # 显式声明所有父类原生参数,和父类__init__签名完全一致 def __init__( self, *, bootstrap=True, ccp_alpha=0.0, class_weight=None, criterion="gini", max_depth=None, max_features="sqrt", max_leaf_nodes=None, max_samples=None, min_impurity_decrease=0.0, min_samples_leaf=1, min_samples_split=2, min_weight_fraction_leaf=0.0, n_estimators=100, n_jobs=None, oob_score=False, random_state=None, verbose=0, warm_start=False, ): # 直接传入固定的n_estimators到父类初始化,自动忽略用户传入的对应参数 super().__init__( bootstrap=bootstrap, ccp_alpha=ccp_alpha, class_weight=class_weight, criterion=criterion, max_depth=max_depth, max_features=max_features, max_leaf_nodes=max_leaf_nodes, max_samples=max_samples, min_impurity_decrease=min_impurity_decrease, min_samples_leaf=min_samples_leaf, min_samples_split=min_samples_split, min_weight_fraction_leaf=min_weight_fraction_leaf, n_estimators=5, n_jobs=n_jobs, oob_score=oob_score, random_state=random_state, verbose=verbose, warm_start=warm_start )
测试验证
fivetrees = FiveTreesClassifier() randomforest = RandomForestClassifier(n_estimators=5) # 两个断言均可正常通过 assert fivetrees.n_estimators == randomforest.n_estimators assert fivetrees.get_params() == randomforest.get_params() # 即使用户主动传入n_estimators也会被强制锁定为5 fivetrees2 = FiveTreesClassifier(n_estimators=100) assert fivetrees2.n_estimators == 5
优缺点
- 优点:完全兼容所有scikit-learn工具,可强制锁定目标参数不被用户传入值、网格搜索等逻辑覆盖,其余参数可正常修改、调参
- 缺点:代码冗余度高,scikit-learn版本升级如果修改了父类的
__init__参数列表,需要手动同步子类的参数签名
方案2:使用
functools.partial快速实现 不需要定义新类,仅需一行代码即可固定参数,适配所有scikit-learn工具。
实现代码
from functools import partial from sklearn.ensemble import RandomForestClassifier FiveTreesClassifier = partial(RandomForestClassifier, n_estimators=5)
测试验证
fivetrees = FiveTreesClassifier() randomforest = RandomForestClassifier(n_estimators=5) # 两个断言均可正常通过 assert fivetrees.n_estimators == randomforest.n_estimators assert fivetrees.get_params() == randomforest.get_params()
优缺点
- 优点:代码极简,无需维护参数签名,scikit-learn版本升级无兼容问题,其余参数可正常修改、调参
- 缺点:无法强制锁定参数,用户主动传入
n_estimators时会覆盖固定值,例如FiveTreesClassifier(n_estimators=100)会得到100棵树的随机森林实例
内容的提问来源于stack exchange,提问作者Enrico Gandini
相关产品推荐
相关产品推荐

