Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 15 additions & 18 deletions .github/PULL_REQUEST_TEMPLATE.md
Original file line number Diff line number Diff line change
@@ -1,24 +1,21 @@
## Purpose
<!-- Thank you for contributing to Tributo! -->
<!-- Please review CONTRIBUTING.md before opening a pull request. -->
<!-- Remove these instructions before submitting your PR. -->

<!-- Clearly describe what this PR does and why. -->
## Description

Closes #
<!-- Briefly describe what this PR changes and why. Mention user-facing,
API, or breaking changes when applicable. -->

## What Changes
## Related issues

<!-- Brief bullet list of changes. -->
<!-- Link related issues, if any. -->

-
## Additional information

## Test Plan

<!-- How did you test? Include commands if applicable. -->

- [ ] 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
<!-- Include when applicable:
- Tests run and their results, including pending or unavailable integration tests
- User-facing, API, or breaking changes and migration notes
- New or changed dependencies and license review status
- Documentation, screenshots, risks, or known limitations
-->
27 changes: 25 additions & 2 deletions src/tributo/explainability/shap.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
76 changes: 76 additions & 0 deletions tests/explainability/test_shap.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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(
Expand Down
Loading