Scikit-learn管道中通过元数据路由向自定义转换器传递参数报错问题
Scikit-learn管道中通过元数据路由向自定义转换器传递参数报错问题
你好,我来帮你排查这个元数据路由的问题~你遇到的报错是因为自定义转换器的元数据路由配置没有正确关联参数到对应的方法,另外代码里还有几个小细节需要调整,我一步步给你说明:
问题根源分析
- 重复定义了
transform方法:你的代码里写了两次def transform(self, X, feature_index_sec=None),第二次会完全覆盖第一次,导致最初的参数校验逻辑失效。 MetadataRouter配置不完整:你只调用了add_self_request,但没有明确指定每个方法需要接收的元数据参数,管道无法知道该把DummyTransformer__feature_index_sec传递到哪个方法里。set_*_request方法实现不够规范:这些方法需要明确标记哪些参数需要被路由,当前的实现没有正确绑定参数请求。
修正后的完整代码
from sklearn.base import BaseEstimator, TransformerMixin from sklearn.pipeline import Pipeline from sklearn.utils.metadata_routing import MetadataRouter, MethodMapping from scipy.sparse import csr_matrix import pandas as pd import numpy as np from sklearn import set_config # Enable metadata routing globally set_config(enable_metadata_routing=True) class DummyTransformer(BaseEstimator, TransformerMixin): # 修复:只保留一个transform方法,保留参数校验逻辑 def transform(self, X, feature_index_sec=None): if feature_index_sec is None: raise ValueError("Missing required argument 'feature_index_sec'.") print(f"Received feature_index_sec with shape: {feature_index_sec.shape}") return X def fit(self, X, y=None, feature_index_sec=None, **fit_params): # 这里可以根据需要使用feature_index_sec,当前是无状态的 return self def fit_transform(self, X, y=None, feature_index_sec=None): self.fit(X, y, feature_index_sec) return self.transform(X, feature_index_sec) def get_metadata_routing(self): print("Configuring metadata routing for DummyTransformer") router = MetadataRouter(owner=self.__class__.__name__) # 明确为每个方法添加参数路由映射 router = router.add_request( feature_index_sec, method_mapping=MethodMapping( fit="feature_index_sec", transform="feature_index_sec", fit_transform="feature_index_sec" ) ) return router # 规范set_*_request方法:明确指定参数是否需要路由 def set_fit_request(self, *, feature_index_sec=True): self._fit_request = {"feature_index_sec": feature_index_sec} return self def set_transform_request(self, *, feature_index_sec=True): self._transform_request = {"feature_index_sec": feature_index_sec} return self def set_fit_transform_request(self, *, feature_index_sec=True): self._fit_transform_request = {"feature_index_sec": feature_index_sec} return self # Dummy data feature_matrix = csr_matrix(np.random.rand(10, 5)) train_idx = pd.DataFrame({'FileDate_ClosingPrice': np.random.rand(10)}) # Configure metadata requests for DummyTransformer transformer = DummyTransformer().set_fit_transform_request(feature_index_sec=True) # Minimal pipeline pipe = Pipeline(steps=[('DummyTransformer', transformer)]) # Test fit_transform pipe.fit_transform(feature_matrix, DummyTransformer__feature_index_sec=train_idx)
关键修改点说明
- 移除重复的
transform方法:保留带有参数校验的那个版本,确保缺失参数时能抛出正确的错误提示。 - 完善
MetadataRouter配置:通过add_request方法把feature_index_sec参数映射到fit、transform、fit_transform三个方法,让管道明确知道这个参数要传递给转换器的哪些方法。 - 规范
set_*_request方法:使用关键字参数(*强制关键字)的方式,明确标记需要路由的参数,这符合scikit-learn元数据路由的规范写法。
运行修正后的代码,你会看到控制台打印出Received feature_index_sec with shape: (10, 1),说明参数已经成功传递到自定义转换器里了。
备注:内容来源于stack exchange,提问作者Jake Drew
相关产品推荐
相关产品推荐

