Skip to content

check_rewind_replaces_observations

yohou.testing.check_rewind_replaces_observations(forecaster, y_reset, X_actual_reset=None, X_future=None, X_forecast=None)

Check rewind() replaces observation buffers correctly.

Parameters

Name Type Description Default
forecaster BaseForecaster

Fitted forecaster instance

required
y_reset DataFrame

New data for reset

required
X_actual_reset DataFrame

Features for reset

None
X_future DataFrame

Known-future features forwarded to rewind()

None
X_forecast DataFrame

External forecast features forwarded to rewind()

None

Raises

Type Description
AssertionError

If observation buffers are not replaced correctly

Source Code

Source code in src/yohou/testing/forecaster.py
def check_rewind_replaces_observations(
    forecaster,
    y_reset: pl.DataFrame,
    X_actual_reset: pl.DataFrame | None = None,
    X_future: pl.DataFrame | None = None,
    X_forecast: pl.DataFrame | None = None,
) -> None:
    """Check rewind() replaces observation buffers correctly.

    Parameters
    ----------
    forecaster : BaseForecaster
        Fitted forecaster instance
    y_reset : pl.DataFrame
        New data for reset
    X_actual_reset : pl.DataFrame, optional
        Features for reset
    X_future : pl.DataFrame, optional
        Known-future features forwarded to rewind()
    X_forecast : pl.DataFrame, optional
        External forecast features forwarded to rewind()

    Raises
    ------
    AssertionError
        If observation buffers are not replaced correctly

    """
    # Precondition: observed_time_ agrees with the observation buffers.
    _assert_observed_time_consistent(forecaster, "rewind")

    # Reset to new data
    forecaster.rewind(y_reset, X_actual_reset, X_future=X_future, X_forecast=X_forecast)

    # Check buffers were replaced
    reset_observed_time = forecaster.observed_time_

    # Handle both panel and non-panel data
    if isinstance(reset_observed_time, dict):
        # Panel data: check each group's observed_time matches
        for group_name in reset_observed_time:
            # All groups share the global 'time' column; the expected reset timestamp is the same for all groups.
            assert reset_observed_time[group_name] == y_reset["time"][-1], (
                f"observed_time_['{group_name}'] should be reset to last time in reset data"
            )
    else:
        # Non-panel data
        assert reset_observed_time == y_reset["time"][-1], "observed_time_ should be reset to last time in reset data"

    if forecaster._y_observed is not None:
        if isinstance(forecaster._y_observed, dict):
            # Panel data
            for group_name, y_obs in forecaster._y_observed.items():
                # _y_observed[group] can be None when observation_horizon == 0
                if y_obs is not None:
                    reset_y_observed_last_time = y_obs["time"][-1]
                    assert reset_y_observed_last_time == reset_observed_time[group_name], (
                        f"Last time in _y_observed['{group_name}'] should match reset observed_time_"
                    )
        else:
            # Non-panel data
            reset_y_observed_last_time = forecaster._y_observed["time"][-1]
            assert reset_y_observed_last_time == reset_observed_time, (
                "Last time in _y_observed should match reset observed_time_ after rewind()"
            )

    if forecaster._X_t_observed is not None:
        if isinstance(forecaster._X_t_observed, dict):
            # Panel data
            for group_name, X_t_obs in forecaster._X_t_observed.items():
                if X_t_obs is not None:
                    reset_X_t_observed_last_time = X_t_obs["time"][-1]
                    assert reset_X_t_observed_last_time == reset_observed_time[group_name], (
                        f"Last time in _X_t_observed['{group_name}'] should match reset observed_time_"
                    )
        else:
            # Non-panel data
            reset_X_t_observed_last_time = forecaster._X_t_observed["time"][-1]
            assert reset_X_t_observed_last_time == reset_observed_time, (
                "Last time in _X_t_observed should match reset observed_time_ after rewind()"
            )