scikit-learn Pipeline会继承最后一个估算器的哪些方法?
scikit-learn Pipeline 方法代理的实际规则
scikit-learn用户指南里说Pipeline“拥有其最后一个估算器的所有方法”,这个表述容易产生误解,实际的运行规则是:
- Pipeline只代理最后一个估算器中属于scikit-learn标准估算器API的方法,比如
predict()、predict_proba()、transform()、score()、decision_function()这类官方定义的方法。 - 它不会自动代理你自定义的任意方法(比如你写的
myfun()),因为Pipeline的设计是围绕scikit-learn的标准接口,而非用户自行添加的方法。
如果要调用自定义估算器的myfun(),你需要直接获取Pipeline里的最后一个估算器实例再调用:
# 方式1:通过步骤名称获取(推荐,更清晰) pipe.named_steps['your_custom_estimator'].myfun() # 方式2:通过索引获取最后一步 pipe[-1].myfun()
背后的逻辑是:Pipeline通过__getattr__方法实现方法代理,但这个方法里有过滤逻辑,只会处理那些属于scikit-learn标准API的方法,而非所有存在于最后估算器上的方法。
内容的提问来源于stack exchange,提问作者Evan Aad
相关产品推荐
相关产品推荐

