diff --git a/.github/PULL_REQUEST_TEMPLATE.md b/.github/PULL_REQUEST_TEMPLATE.md index 84bab99..c572880 100644 --- a/.github/PULL_REQUEST_TEMPLATE.md +++ b/.github/PULL_REQUEST_TEMPLATE.md @@ -1,24 +1,21 @@ -## Purpose + + + - +## Description -Closes # + -## What Changes +## Related issues - + -- +## Additional information -## Test Plan - - - -- [ ] Unit tests pass -- [ ] Lint passes (`ruff check .`) -- [ ] Integration tests pass (if applicable) - -## Open Source Checklist - -- [ ] No internal credentials, URLs, or tokens exposed -- [ ] New external dependencies reviewed for license compatibility + diff --git a/src/tributo/explainability/shap.py b/src/tributo/explainability/shap.py index d78f3ec..804af9b 100644 --- a/src/tributo/explainability/shap.py +++ b/src/tributo/explainability/shap.py @@ -269,7 +269,7 @@ def explain_batch( data = getattr(explanation, "data", batch) data_array = np.asarray(data) base_values = _normalise_base_values( - getattr(explanation, "base_values", None), values + getattr(explanation, "base_values", None), values, labels=labels ) if request.output_target == "log_loss": model_outputs = values.sum(axis=1) + base_values @@ -394,10 +394,33 @@ def _tree_model_output(output_target: str) -> str | None: }.get(output_target, output_target) -def _normalise_base_values(raw: Any, values: np.ndarray) -> np.ndarray: +def _normalise_base_values( + raw: Any, + values: np.ndarray, + *, + labels: np.ndarray | None = None, +) -> np.ndarray: if raw is None: return np.zeros((values.shape[0], values.shape[2]), dtype=np.float64) + if callable(raw): + if labels is None: + raise ValueError("callable SHAP base_values require labels") + raw = np.asarray([raw(label) for label in labels]) array = np.asarray(raw) + if array.dtype == object: + dynamic_values = array.reshape(-1) + if any(callable(value) for value in dynamic_values): + if labels is None: + raise ValueError("callable SHAP base_values require labels") + if not all(callable(value) for value in dynamic_values): + raise ValueError("SHAP base_values contain mixed callable values") + if array.ndim != 1 or array.shape[0] != values.shape[0]: + raise ValueError( + "callable SHAP base_values must contain one value per sample" + ) + array = np.asarray( + [value(label) for value, label in zip(array, labels, strict=True)] + ) if array.ndim == 0: return np.full((values.shape[0], values.shape[2]), float(array)) if array.ndim == 1: diff --git a/tests/explainability/test_shap.py b/tests/explainability/test_shap.py index 6497f66..3fb064e 100644 --- a/tests/explainability/test_shap.py +++ b/tests/explainability/test_shap.py @@ -19,6 +19,28 @@ class _Explanation: base_values = np.asarray([0.5, 0.6]) +class _DynamicExplanation: + values = np.asarray([[[0.1], [1.0]], [[0.1], [2.0]]]) + data = np.asarray([[10.0, 20.0], [30.0, 40.0]]) + + def base_values(self, label: float) -> float: + return 0.5 if label == 0 else 0.6 + + +class _ArrayDynamicExplanation: + values = np.asarray([[[0.1], [1.0]], [[0.1], [2.0]]]) + data = np.asarray([[10.0, 20.0], [30.0, 40.0]]) + + def __init__(self) -> None: + self.base_values = np.asarray( + [self._base_value, self._base_value], dtype=object + ) + + @staticmethod + def _base_value(label: float) -> float: + return 0.5 if label == 0 else 0.6 + + def test_shap_long_rows_use_top_k_and_preserve_provenance() -> None: request = _request(limits={"top_k": 1}) prepared = PreparedExplainer( @@ -65,6 +87,60 @@ def test_tree_log_loss_requires_labels_at_adapter_boundary() -> None: ) +def test_tree_log_loss_materialises_dynamic_base_values() -> None: + request = _request( + backend="tree", + output_target="log_loss", + label_column="label", + reference={"uri": "/data/reference.npy"}, + ) + prepared = PreparedExplainer( + backend="tree", + exactness="exact", + explain=lambda batch, **kwargs: _DynamicExplanation(), + feature_names=("feature_a", "feature_b"), + ) + + rows = ShapAdapter().explain_batch( + prepared, + np.asarray([[10.0, 20.0], [30.0, 40.0]], dtype=np.float32), + input_ids=("row-1", "row-2"), + model_digest="a" * 64, + request=request, + labels=np.asarray([0, 1]), + ) + + assert [row.base_value for row in rows[::2]] == [0.5, 0.6] + assert [row.model_output for row in rows[::2]] == [1.6, 2.7] + + +def test_tree_log_loss_materialises_array_dynamic_base_values() -> None: + request = _request( + backend="tree", + output_target="log_loss", + label_column="label", + reference={"uri": "/data/reference.npy"}, + ) + prepared = PreparedExplainer( + backend="tree", + exactness="exact", + explain=lambda batch, **kwargs: _ArrayDynamicExplanation(), + feature_names=("feature_a", "feature_b"), + ) + + rows = ShapAdapter().explain_batch( + prepared, + np.asarray([[10.0, 20.0], [30.0, 40.0]], dtype=np.float32), + input_ids=("row-1", "row-2"), + model_digest="a" * 64, + request=request, + labels=np.asarray([0, 1]), + ) + + assert [row.base_value for row in rows[::2]] == [0.5, 0.6] + assert [row.model_output for row in rows[::2]] == [1.6, 2.7] + + def test_sensitive_feature_values_are_opt_in() -> None: request = _request(result_policy={"allow_sensitive_features": True}) prepared = PreparedExplainer(