Check probabilities sum to 1.0 per row per target (tolerance: 1e-6).
Parameters
| Name |
Type |
Description |
Default |
forecaster
|
BaseClassProbaForecaster
|
Fitted class-probability forecaster instance.
|
required
|
y_test
|
DataFrame
|
|
required
|
Raises
| Type |
Description |
AssertionError
|
If probabilities do not sum to approximately 1.0.
|
Source Code
View on GitHub
Source code in src/yohou/testing/class_proba.py
| def check_class_proba_prediction_sums(forecaster, y_test: pl.DataFrame) -> None:
"""Check probabilities sum to 1.0 per row per target (tolerance: 1e-6).
Parameters
----------
forecaster : BaseClassProbaForecaster
Fitted class-probability forecaster instance.
y_test : pl.DataFrame
Test target data.
Raises
------
AssertionError
If probabilities do not sum to approximately 1.0.
"""
forecasting_horizon = min(3, len(y_test))
y_pred = forecaster.predict_class_proba(forecasting_horizon=forecasting_horizon)
_, y_panel_groups = inspect_panel(y_test)
if len(y_panel_groups) > 0:
for group_prefix in y_panel_groups:
for target_col, class_labels in forecaster.classes_.items():
proba_cols = [f"{group_prefix}__{target_col}_proba_{label}" for label in class_labels]
row_sums = y_pred.select(proba_cols).sum_horizontal()
max_err = (row_sums - 1.0).abs().max()
assert max_err < 1e-6, (
f"Probabilities for {group_prefix}__{target_col} deviate from 1.0 by up to {max_err:.2e}, expected < 1e-6"
)
else:
for target_col, class_labels in forecaster.classes_.items():
proba_cols = [f"{target_col}_proba_{label}" for label in class_labels]
row_sums = y_pred.select(proba_cols).sum_horizontal()
max_err = (row_sums - 1.0).abs().max()
assert max_err < 1e-6, (
f"Probabilities for {target_col} deviate from 1.0 by up to {max_err:.2e}, expected < 1e-6"
)
|