Skip to content

check_class_proba_prediction_bounds

yohou.testing.check_class_proba_prediction_bounds(forecaster, y_test)

Check all probability values are in [0, 1].

Parameters

Name Type Description Default
forecaster BaseClassProbaForecaster

Fitted class-probability forecaster instance.

required
y_test DataFrame

Test target data.

required

Raises

Type Description
AssertionError

If any probability value is outside [0, 1].

Source Code

Source code in src/yohou/testing/class_proba.py
def check_class_proba_prediction_bounds(forecaster, y_test: pl.DataFrame) -> None:
    """Check all probability values are in [0, 1].

    Parameters
    ----------
    forecaster : BaseClassProbaForecaster
        Fitted class-probability forecaster instance.
    y_test : pl.DataFrame
        Test target data.

    Raises
    ------
    AssertionError
        If any probability value is outside [0, 1].

    """
    forecasting_horizon = min(3, len(y_test))
    y_pred = forecaster.predict_class_proba(forecasting_horizon=forecasting_horizon)

    proba_cols = [col for col in y_pred.columns if "_proba_" in col]
    assert len(proba_cols) > 0, "No probability columns found"

    mins = y_pred.select(proba_cols).min().row(0)
    maxs = y_pred.select(proba_cols).max().row(0)
    for col, min_val, max_val in zip(proba_cols, mins, maxs, strict=True):
        assert min_val >= 0.0, f"Probability column {col} has negative values (min={min_val})"
        assert max_val <= 1.0, f"Probability column {col} has values > 1 (max={max_val})"