Skip to content

check_predict_X_forecast_override

yohou.testing.check_predict_X_forecast_override(forecaster, X_forecast, forecasting_horizon=3)

Check predict with X_forecast override produces different results.

Validates that passing X_forecast at predict time overrides the stored forecasts without mutating forecaster state.

Parameters

Name Type Description Default
forecaster BaseForecaster

Fitted forecaster instance (fitted with X_forecast).

required
X_forecast DataFrame

External forecasts for override.

required
forecasting_horizon int

Number of steps ahead to forecast.

3

Source Code

Source code in src/yohou/testing/forecaster.py
def check_predict_X_forecast_override(
    forecaster,
    X_forecast: pl.DataFrame,
    forecasting_horizon: int = 3,
) -> None:
    """Check predict with X_forecast override produces different results.

    Validates that passing X_forecast at predict time overrides the stored
    forecasts without mutating forecaster state.

    Parameters
    ----------
    forecaster : BaseForecaster
        Fitted forecaster instance (fitted with X_forecast).
    X_forecast : pl.DataFrame
        External forecasts for override.
    forecasting_horizon : int, default=3
        Number of steps ahead to forecast.

    """
    # Store a snapshot of the original raw so we can detect in-place mutation.
    original_raw = forecaster._X_forecast_raw_
    if original_raw is not None:
        original_raw = original_raw.clone()

    # Predict with override
    y_pred = forecaster.predict(
        forecasting_horizon=forecasting_horizon,
        X_forecast=X_forecast,
    )

    assert isinstance(y_pred, pl.DataFrame), (
        f"predict() with X_forecast override must return pl.DataFrame, got {type(y_pred).__name__}"
    )

    # State unchanged (predict must not mutate stored raw)
    if original_raw is not None:
        assert forecaster._X_forecast_raw_.equals(original_raw), (
            "predict() with X_forecast override must not mutate _X_forecast_raw_"
        )