You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用ForecastingGridSearchCV进行网格搜索时触发KeyError: 'estimator'错误的求助

ForecastingGridSearchCV进行网格搜索时触发KeyError: 'estimator'错误的求助

我在尝试对目标变量y和特征变量X进行可选缩放后拟合预测模型,但在运行网格搜索时遇到了KeyError: 'estimator'的错误,以下是我的代码和错误详情:

导入的库

from sklearn.preprocessing import MinMaxScaler, PowerTransformer, RobustScaler
from sktime.performance_metrics.forecasting import MeanSquaredError
from sktime.forecasting.statsforecast import StatsForecastAutoARIMA
from sktime.forecasting.model_selection import ExpandingWindowSplitter
from sktime.forecasting.compose import ForecastingPipeline
from sktime.forecasting.fbprophet import Prophet

数据集

训练集train1

y   M0  Delta
ds          
2022-10 127 249.0   14.0
2022-11 194 298.0   160.0
2022-12 202 269.0   128.0
2023-01 149 215.0   33.0
2023-02 203 198.0   111.0
2023-03 222 259.0   10.0
2023-04 195 232.0   -1.0
2023-05 236 210.0   47.0
2023-06 76  155.0   -12.0
2023-07 204 232.0   101.0
2023-08 231 221.0   -13.0
2023-09 171 196.0   -30.0
2023-10 172 80.0    39.0
2023-11 173 206.0   1.0
2023-12 132 189.0   -8.0
2024-01 165 190.0   50.0
2024-02 93  136.0   0.0
2024-03 105 126.0   45.0
2024-04 99  134.0   18.0
2024-05 121 128.0   12.0
2024-06 109 181.0   0.0
2024-07 183 244.0   7.0
2024-08 159 195.0   27.0
2024-09 147 186.0   -5.0

测试集test1

M0  Delta
ds      
2024-10 161.0   -10.0

预测管道与网格搜索代码

pipe_y = TransformedTargetForecaster(
        steps=[
            ("scaler", OptionalPassthrough(TabularToSeriesAdaptor(RobustScaler()))),
            ("forecaster", StatsForecastAutoARIMA(sp=12)),
        ]
    )
pipe_X = ForecastingPipeline(
    steps=[
        ("scaler", OptionalPassthrough(TabularToSeriesAdaptor(RobustScaler()))),
        ("forecaster", pipe_y),
    ]
)

cv1 = ExpandingWindowSplitter(fh=[1], initial_window=train1.shape[0]-3, step_length=1)

gscv1 = ForecastingGridSearchCV(
    forecaster=pipe_X,
    param_grid=[
        {
            "scaler__passthrough": [True, False],
            "forecaster__scaler__passthrough": [True, False],
            "forecaster": [StatsForecastAutoARIMA(sp=12)],
            },
        {
            "scaler__passthrough": [True, False],
            "forecaster": [Prophet()],
            "forecaster__scaler__passthrough": [True, False],
            "forecaster__seasonality_mode": ['addictive','multiplicative'],
            "forecaster__changepoint_prior_scale": [0.001, 0.01, 0.1, 0.5],
            "forecaster__seasonality_prior_scale": [0.01, 0.1, 1.0, 10.0],
            },
        ],
    cv=cv1,
    error_score="raise",
    scoring=MeanSquaredError(square_root=True),
)

gscv1.fit(train1['y'], X=train1[features1])

错误回溯信息

---------------------------------------------------------------------------
_RemoteTraceback                          Traceback (most recent call last)
_RemoteTraceback: 
"""
Traceback (most recent call last):
  File "c:\Desktop\Waterfall Chart\lib\site-packages\joblib\externals\loky\process_executor.py", line 463, in _process_worker
    r = call_item()
  File "c:\Desktop\Waterfall Chart\lib\site-packages\joblib\externals\loky\process_executor.py", line 291, in __call__
    return self.fn(*self.args, **self.kwargs)
  File "c:\Desktop\Waterfall Chart\lib\site-packages\joblib\parallel.py", line 598, in __call__
    return [func(*args, **kwargs)
  File "c:\Desktop\Waterfall Chart\lib\site-packages\joblib\parallel.py", line 598, in <listcomp>
    return [func(*args, **kwargs)
  File "c:\Desktop\Waterfall Chart\lib\site-packages\sktime\forecasting\model_selection\_tune.py", line 368, in _fit_and_score
    forecaster.set_params(**params)
  File "c:\Desktop\Waterfall Chart\lib\site-packages\sktime\base\_meta.py", line 60, in set_params
    self._set_params(steps_attr, **kwargs)
  File "c:\Desktop\Waterfall Chart\lib\site-packages\sktime\base\_meta.py", line 142, in _set_params
    super().set_params(**params)
  File "c:\Desktop\Waterfall Chart\lib\site-packages\skbase\base\_base.py", line 398, in set_params
    unmatched_params = {key: params[key] for key in unmatched_keys}
  File "c:\Desktop\Waterfall Chart\lib\site-packages\skbase\base\_base.py", line 398, in <dictcomp>
    unmatched_params = {key: params[key] for key in unmatched_keys}
KeyError: 'estimator'
"""
...
--> 763         raise self._result
    764     return self._result
    765 finally:

KeyError: 'estimator'

问题分析与解决方案

这个错误的核心原因是参数网格中直接替换forecaster时,破坏了原管道的嵌套结构,导致参数路径失效:

  1. 你的pipe_X中的forecaster是pipe_y(一个TransformedTargetForecaster),它包含scaler和forecaster两个步骤;但在参数网格中,你直接把forecaster替换成了StatsForecastAutoARIMA或Prophet(单个预测器),此时forecaster__scaler__passthrough这个参数路径就不存在了,因为单个预测器没有scaler子步骤,进而触发参数匹配错误。

  2. 另外,OptionalPassthrough的参数设置和嵌套管道的参数传递逻辑也可能导致内部的estimator参数找不到。

修复步骤

1. 统一预测器的嵌套结构

不管是用StatsForecastAutoARIMA还是Prophet,都要把它们包裹在TransformedTargetForecaster中,保持和原pipe_y一致的结构:

# 先定义带可选缩放的基础预测器模板
def create_target_scaled_forecaster(base_forecaster):
    return TransformedTargetForecaster(
        steps=[
            ("scaler", OptionalPassthrough(TabularToSeriesAdaptor(RobustScaler()))),
            ("forecaster", base_forecaster),
        ]
    )

2. 调整参数网格的结构

针对不同的预测器,分别定义对应的参数,确保参数路径和管道结构匹配:

gscv1 = ForecastingGridSearchCV(
    forecaster=pipe_X,
    param_grid=[
        # 针对StatsForecastAutoARIMA的参数组
        {
            "scaler__passthrough": [True, False],
            "forecaster": [create_target_scaled_forecaster(StatsForecastAutoARIMA(sp=12))],
            "forecaster__scaler__passthrough": [True, False],
        },
        # 针对Prophet的参数组
        {
            "scaler__passthrough": [True, False],
            "forecaster": [create_target_scaled_forecaster(Prophet())],
            "forecaster__scaler__passthrough": [True, False],
            "forecaster__forecaster__seasonality_mode": ['additive','multiplicative'],  # 注意参数路径:外层是TransformedTargetForecaster,内部才是Prophet
            "forecaster__forecaster__changepoint_prior_scale": [0.001, 0.01, 0.1, 0.5],
            "forecaster__forecaster__seasonality_prior_scale": [0.01, 0.1, 1.0, 10.0],
        },
    ],
    cv=cv1,
    error_score="raise",
    scoring=MeanSquaredError(square_root=True),
)

3. 验证管道结构(可选)

可以先打印管道的所有可设置参数,确认参数路径是否正确:

print(pipe_X.get_params().keys())

这样就能看到所有合法的参数名称,确保网格搜索中的参数路径和这些名称一致。


备注:内容来源于stack exchange,提问作者SM9595

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.14 14:34:51