Skip to content

check_observe_extends_observations

yohou.testing.check_observe_extends_observations(forecaster, y_observe, X_actual_observe=None, X_future=None, X_forecast=None)

Check observe() extends observation buffers correctly.

Parameters

Name Type Description Default
forecaster BaseForecaster

Fitted forecaster instance

required
y_observe DataFrame

New data for update

required
X_actual_observe DataFrame

Features for update

None
X_future DataFrame

Known-future features forwarded to observe()

None
X_forecast DataFrame

External forecast features forwarded to observe()

None

Raises

Type Description
AssertionError

If observation buffers are not extended correctly

Source Code

Source code in src/yohou/testing/forecaster.py
def check_observe_extends_observations(
    forecaster,
    y_observe: pl.DataFrame,
    X_actual_observe: pl.DataFrame | None = None,
    X_future: pl.DataFrame | None = None,
    X_forecast: pl.DataFrame | None = None,
) -> None:
    """Check observe() extends observation buffers correctly.

    Parameters
    ----------
    forecaster : BaseForecaster
        Fitted forecaster instance
    y_observe : pl.DataFrame
        New data for update
    X_actual_observe : pl.DataFrame, optional
        Features for update
    X_future : pl.DataFrame, optional
        Known-future features forwarded to observe()
    X_forecast : pl.DataFrame, optional
        External forecast features forwarded to observe()

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

    """
    # Store original buffer length
    original_observed_time = forecaster.observed_time_

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

    # Update with new data
    forecaster.observe(y_observe, X_actual_observe, X_future=X_future, X_forecast=X_forecast)

    # Check buffers were extended
    updated_observed_time = forecaster.observed_time_

    # Handle both panel and non-panel data for comparison
    if isinstance(updated_observed_time, dict):
        # Panel data: check all groups were updated
        for group_name in updated_observed_time:
            assert updated_observed_time[group_name] >= original_observed_time[group_name], (
                f"observed_time_ for group {group_name} should be updated"
            )
    else:
        # Non-panel data
        assert updated_observed_time >= original_observed_time, (
            "observed_time_ should be updated to at least the last time in update 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:
                    updated_y_observed_last_time = y_obs["time"][-1]
                    assert updated_y_observed_last_time == updated_observed_time[group_name], (
                        f"Last time in _y_observed['{group_name}'] should match updated observed_time_"
                    )
        else:
            # Non-panel data
            updated_y_observed_last_time = forecaster._y_observed["time"][-1]
            assert updated_y_observed_last_time == updated_observed_time, (
                "Last time in _y_observed should match updated observed_time_ after observe()"
            )

    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:
                    updated_X_t_observed_last_time = X_t_obs["time"][-1]
                    assert updated_X_t_observed_last_time == updated_observed_time[group_name], (
                        f"Last time in _X_t_observed['{group_name}'] should match updated observed_time_"
                    )
        else:
            # Non-panel data
            updated_X_t_observed_last_time = forecaster._X_t_observed["time"][-1]
            assert updated_X_t_observed_last_time == updated_observed_time, (
                "Last time in _X_t_observed should match updated observed_time_ after observe()"
            )