File size: 1,594 Bytes
2532605 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 | """
Tests for prediction helpers.
Run: python -m pytest tests/test_predict.py -v
"""
import sys
from types import SimpleNamespace
import numpy as np
from ml.feature_engineering import prepare_model_data
from ml.predict import predict_dataframe
from tests.test_feature_engineering import sample_race_df # noqa: F401
class DummyModel:
def predict_proba(self, X):
return np.column_stack([np.full(len(X), 0.75), np.full(len(X), 0.25)])
class DummyExplainer:
def __init__(self, model):
self.model = model
def shap_values(self, X):
return np.ones((len(X), len(X.columns)))
def test_predict_dataframe_returns_shap_dicts(sample_race_df, monkeypatch):
_, encoders = prepare_model_data(sample_race_df)
dummy_shap = SimpleNamespace(TreeExplainer=DummyExplainer)
monkeypatch.setitem(sys.modules, "shap", dummy_shap)
predictions = predict_dataframe(
sample_race_df.iloc[:2], DummyModel(), encoders, explain=True
)
assert "shap_values" in predictions.columns
assert predictions["win_probability"].tolist() == [0.25, 0.25]
assert isinstance(predictions.loc[0, "shap_values"], dict)
assert set(predictions.loc[0, "shap_values"]) == set(
prepare_model_data(sample_race_df.iloc[:2], encoders=encoders)[0].X.columns
)
def test_predict_dataframe_can_skip_explanations(sample_race_df):
_, encoders = prepare_model_data(sample_race_df)
predictions = predict_dataframe(
sample_race_df.iloc[:2], DummyModel(), encoders, explain=False
)
assert "shap_values" not in predictions.columns
|