diff --git a/docs/how-to/explainability.md b/docs/how-to/explainability.md index 0c6aec4..6f49ac7 100644 --- a/docs/how-to/explainability.md +++ b/docs/how-to/explainability.md @@ -86,6 +86,35 @@ operation whose lease has expired can be reclaimed only with `force_resume=true`, after confirming that the previous driver is no longer active. +## Select XGBoost outputs + +For an XGBoost native `gbtree` or `dart` model, `model_output`, `raw`, and +`raw_margin` use the Booster contribution API with exact TreeSHAP and the +strict output shapes supported by XGBoost 2.1 or newer. The final contribution +column is recorded as the base value, and Tributo verifies that the base value +plus feature contributions reconstructs each raw model output. Linear boosters +are rejected because their native contributions are not TreeSHAP. Probability, +log-loss, and requests with reference data continue to use SHAP TreeExplainer. + +Multi-class requests explain every class by default. The request's +`output_target` must match the value declared by the Bundle descriptor. With the +default export configuration above, set `output_selection` to `predicted` to +retain only the class selected by the raw model margins: + +```json +{ + "output_target": "model_output", + "output_selection": "predicted" +} +``` + +The output keeps the original class index, such as `output_7`; it is not +renumbered after selection. `limits.top_k` is evaluated against that selected +class. Binary classifiers have one margin contribution group, so `predicted` +and `all` produce the same output space. The `predicted` policy is not accepted +for regression, probability, log-loss, requests with reference data, or +model-agnostic requests. + ## Read results Results are written as sharded Parquet in long format. `result_uri` in the @@ -96,7 +125,8 @@ another attempt's files. Each row identifies an input, output, and feature and includes the contribution, base value, output semantics, backend, exactness, model digest, and optional preprocessor/feature map digests. `receipt.json` records the result digest, row and byte counts, reference provenance, -dependency versions, and the declared access/privacy/retention policy. +dependency versions, output selection, and the declared +access/privacy/retention policy. Consumers should read the `result_uri` and `receipt_uri` from the persisted operation record, or use the returned receipt, rather than assuming that the diff --git a/src/tributo/explainability/contracts.py b/src/tributo/explainability/contracts.py index a377c2c..1c14dca 100644 --- a/src/tributo/explainability/contracts.py +++ b/src/tributo/explainability/contracts.py @@ -247,6 +247,7 @@ class ExplainabilityRequest(_FrozenContract): backend: Literal["auto", "tree", "model_agnostic", "deep", "gradient"] = "auto" feature_view: Literal["raw", "transformed", "model_input"] = "raw" output_target: str = Field(default="model_output", min_length=1) + output_selection: Literal["all", "predicted"] = "all" label_column: str | None = Field(default=None, min_length=1) allow_approximate: bool = False reference: ReferenceBinding | None = None @@ -358,6 +359,7 @@ class ExplainabilityReceipt(_FrozenContract): exactness: Literal["exact", "approximate", "conditional"] feature_view: Literal["raw", "transformed", "model_input"] output_target: str = Field(min_length=1) + output_selection: Literal["all", "predicted"] = "all" execution_profile: str = Field(default="batch", min_length=1) input_rows: int = Field(default=0, ge=0) explanation_rows: int = Field(default=0, ge=0) diff --git a/src/tributo/explainability/executor.py b/src/tributo/explainability/executor.py index bc88c54..98fad71 100644 --- a/src/tributo/explainability/executor.py +++ b/src/tributo/explainability/executor.py @@ -217,6 +217,7 @@ def run_batch_explainability( result_bytes = 0 try: _validate_request_against_descriptor(manifest, request) + output_count = _explanation_output_count_upper_bound(manifest, request) selection = resolver.describe(request.input) opened = resolver.open(selection) dataset = opened.dataset @@ -227,7 +228,11 @@ def run_batch_explainability( f"input rows {input_rows} exceed limits.max_rows=" f"{request.limits.max_rows}" ) - ExplainabilityPlanner.preflight_limits(request, input_rows=input_rows) + ExplainabilityPlanner.preflight_limits( + request, + input_rows=input_rows, + output_count=output_count, + ) explained = dataset.map_batches( worker, @@ -417,6 +422,7 @@ def _make_receipt( exactness=exactness, feature_view=request.feature_view, output_target=request.output_target, + output_selection=request.output_selection, input_rows=input_rows, explanation_rows=explanation_rows, result_uri=result_uri, @@ -458,7 +464,14 @@ def __init__( self._runtime: BundleModelRuntime | None = None self._artifact_stack = ExitStack() self._context = self._load_context(request) - plan = ExplainabilityPlanner(registry).plan(self._context, request) + plan = ExplainabilityPlanner(registry).plan( + self._context, + request, + output_count=_explanation_output_count_upper_bound( + self._manifest, + request, + ), + ) self._plan = plan self._prepared = plan.adapter().prepare(self._context, request) @@ -943,6 +956,49 @@ def _selected_artifact(manifest: Any, request: ExplainabilityRequest) -> Any: ) from exc +def _explanation_output_count_upper_bound( + manifest: Any, + request: ExplainabilityRequest, +) -> int: + """Resolve a safe attribution-output bound from a verified manifest.""" + artifact = _selected_artifact(manifest, request) + if artifact.flavor_id != "xgboost-native-v1": + return 1 if request.output_target in {"raw", "raw_margin"} else 2 + + signature = getattr(manifest, "output_signature", None) + fields = tuple(getattr(signature, "output_fields", ())) + probability_fields = tuple( + field + for field in fields + if str(getattr(field, "name", "")).lower() + in {"probability", "probabilities", "proba", "scores"} + ) + prediction_fields = tuple( + field + for field in fields + if str(getattr(field, "name", "")).lower() in {"prediction", "predictions"} + ) + task_type = getattr(getattr(manifest, "source_info", None), "task_type", None) + if task_type == "regression": + candidates = prediction_fields + elif task_type == "classification": + candidates = probability_fields + else: + candidates = probability_fields or prediction_fields + if len(candidates) != 1: + raise ValueError( + "XGBoost native explainability requires one typed probability or " + "prediction output signature" + ) + shape = tuple(getattr(candidates[0], "shape", ())) + if len(shape) != 2 or not isinstance(shape[1], int) or shape[1] < 1: + raise ValueError( + "XGBoost native explainability requires a fixed output dimension " + "in the typed manifest signature" + ) + return shape[1] + + def _model_digest(manifest: Any, request: ExplainabilityRequest) -> str: return str(_selected_artifact(manifest, request).tree_digest) diff --git a/src/tributo/explainability/planner.py b/src/tributo/explainability/planner.py index 97b7961..ae615e2 100644 --- a/src/tributo/explainability/planner.py +++ b/src/tributo/explainability/planner.py @@ -38,21 +38,28 @@ def __init__(self, registry: ExplainerRegistry | None = None) -> None: @staticmethod def preflight_limits( - request: ExplainabilityRequest, *, input_rows: int + request: ExplainabilityRequest, + *, + input_rows: int, + output_count: int, ) -> dict[str, int | float]: """Reject a request whose known upper bound exceeds its byte budget.""" + if output_count < 1: + raise ValueError("output_count must be positive") limit = request.limits.max_explanation_bytes feature_count = min( len(request.feature_columns) or request.limits.max_features or 1, request.limits.max_features or len(request.feature_columns) or 1, ) - output_count = 1 if request.output_target in {"raw", "raw_margin"} else 2 + effective_output_count = ( + 1 if request.output_selection == "predicted" else output_count + ) background_rows = ( request.limits.max_background_rows or (request.reference.rows if request.reference is not None else None) or 1 ) - estimated_rows = input_rows * feature_count * output_count + estimated_rows = input_rows * feature_count * effective_output_count estimated_bytes = estimated_rows * 512 if limit is not None and estimated_bytes > limit: raise ValueError( @@ -70,6 +77,7 @@ def preflight_limits( return { "estimated_output_rows": estimated_rows, "estimated_output_bytes": estimated_bytes, + "estimated_output_count": effective_output_count, "estimated_background_rows": background_rows, "batch_size": request.resource_policy.batch_size, "concurrency": request.resource_policy.concurrency, @@ -79,6 +87,8 @@ def plan( self, context: ExplainableModelContext, request: ExplainabilityRequest, + *, + output_count: int, ) -> ExplainabilityPlan: adapter_id = f"{request.explainer}-v1" adapter = self._registry.get(adapter_id) @@ -101,6 +111,7 @@ def plan( resource_requirements=self.preflight_limits( request, input_rows=0, + output_count=output_count, ), ) diff --git a/src/tributo/explainability/shap.py b/src/tributo/explainability/shap.py index 804af9b..67226a0 100644 --- a/src/tributo/explainability/shap.py +++ b/src/tributo/explainability/shap.py @@ -3,6 +3,8 @@ from __future__ import annotations import importlib.util +import json +from dataclasses import dataclass from typing import Any, ClassVar import numpy as np @@ -19,6 +21,103 @@ ) from tributo.util.annotations import PublicAPI +_NATIVE_OUTPUT_TARGETS = frozenset({"model_output", "raw", "raw_margin"}) + + +@dataclass(frozen=True) +class _NativeExplanation: + values: np.ndarray + data: np.ndarray + base_values: np.ndarray + model_outputs: np.ndarray + + +class _NativeTreeExplainer: + """Strict XGBoost contribution wrapper with a SHAP-like result shape.""" + + def __init__( + self, + booster: Any, + *, + feature_names: tuple[str, ...], + objective: str | None, + ) -> None: + self._booster = booster + self._feature_names = feature_names + self._objective = objective + + def __call__( + self, + values: np.ndarray, + *, + check_additivity: bool = False, + ) -> _NativeExplanation: + del check_additivity + import xgboost + + data = np.asarray(values, dtype=np.float32) + if data.ndim != 2: + raise ValueError( + f"XGBoost native TreeSHAP input must be 2-D, got {data.shape}" + ) + if self._feature_names and len(self._feature_names) != data.shape[1]: + raise ValueError( + "XGBoost feature names count does not match the explanation input" + ) + feature_types = tuple(self._booster.feature_types or ()) + if feature_types and len(feature_types) != data.shape[1]: + raise ValueError( + "XGBoost feature types count does not match the explanation input" + ) + matrix = xgboost.DMatrix( + data, + feature_names=list(self._feature_names) or None, + feature_types=list(feature_types) or None, + ) + contributions = np.asarray( + self._booster.predict( + matrix, + pred_contribs=True, + approx_contribs=False, + strict_shape=True, + ) + ) + model_outputs = np.asarray( + self._booster.predict( + matrix, + output_margin=True, + strict_shape=True, + ) + ) + if model_outputs.ndim != 2 or model_outputs.shape[0] != data.shape[0]: + raise ValueError( + "XGBoost raw outputs violate the strict shape contract: " + f"{model_outputs.shape}" + ) + output_count = model_outputs.shape[1] + if self._objective and self._objective.startswith("multi:"): + if output_count < 2: + raise ValueError( + "XGBoost multiclass raw outputs must contain multiple groups" + ) + elif self._objective and _supported_objective(self._objective): + if output_count != 1: + raise ValueError( + "XGBoost regression/binary raw outputs must contain one group" + ) + expected_shape = (data.shape[0], output_count, data.shape[1] + 1) + if contributions.shape != expected_shape: + raise ValueError( + "XGBoost contributions violate the strict shape contract: " + f"expected {expected_shape}, got {contributions.shape}" + ) + return _NativeExplanation( + values=np.transpose(contributions[:, :, :-1], (0, 2, 1)), + data=data, + base_values=contributions[:, :, -1], + model_outputs=model_outputs, + ) + @PublicAPI(stability="alpha") class ShapAdapter: @@ -34,12 +133,14 @@ def supports( context: ExplainableModelContext, request: ExplainabilityRequest, ) -> SupportDecision: - missing = _missing("shap") if ( request.backend in ("auto", "tree") and context.flavor_id == "xgboost-native-v1" ): - missing += _missing("xgboost") + native_strategy = _uses_native_tree_strategy(context, request) + missing = _missing("xgboost") + if not native_strategy: + missing += _missing("shap") if context.artifact_format not in {"ubj", "xgboost-json"}: return SupportDecision( supported=False, @@ -71,6 +172,29 @@ def supports( backend="tree", exactness="exact", ) + if request.output_selection == "predicted": + if not native_strategy: + return SupportDecision( + supported=False, + reason=( + "output_selection='predicted' is only supported for " + "native XGBoost raw model outputs" + ), + backend="tree", + exactness="exact", + ) + if not context.objective or not context.objective.startswith( + ("binary:", "multi:") + ): + return SupportDecision( + supported=False, + reason=( + "output_selection='predicted' requires an XGBoost " + "classification objective" + ), + backend="tree", + exactness="exact", + ) if ( request.output_target in {"probability", "log_loss"} and request.reference is None @@ -106,6 +230,17 @@ def supports( ) if request.backend in ("auto", "model_agnostic"): + missing = _missing("shap") + if request.output_selection == "predicted": + return SupportDecision( + supported=False, + reason=( + "output_selection='predicted' is not implemented for " + "model-agnostic SHAP" + ), + backend="model_agnostic", + exactness="approximate", + ) if request.output_target == "log_loss": return SupportDecision( supported=False, @@ -171,7 +306,6 @@ def prepare( decision = self.supports(context, request) if not decision.supported: raise ValueError(decision.reason) - shap = _require_shap() feature_names = context.feature_names if decision.backend == "tree": booster = context.model_object @@ -182,6 +316,28 @@ def prepare( booster = xgboost.Booster() booster.load_model(str(context.artifact_path)) + if not feature_names: + feature_names = tuple(booster.feature_names or ()) + if _uses_native_tree_strategy(context, request): + booster_kind = _xgboost_booster_kind(booster) + if booster_kind not in {"gbtree", "dart"}: + raise ValueError( + "XGBoost native TreeSHAP requires a gbtree or dart " + f"booster, got {booster_kind!r}" + ) + return PreparedExplainer( + backend="tree", + exactness="exact", + explain=_NativeTreeExplainer( + booster, + feature_names=feature_names, + objective=context.objective, + ), + feature_names=feature_names, + preprocessor_digest=context.preprocessor_digest, + feature_map_digest=context.feature_map_digest, + ) + shap = _require_shap() model_output = _tree_model_output(request.output_target) kwargs: dict[str, Any] = {} if model_output is not None: @@ -191,8 +347,6 @@ def prepare( kwargs["data"] = np.asarray(reference_data) kwargs["feature_perturbation"] = "interventional" explainer = shap.TreeExplainer(booster, **kwargs) - if not feature_names: - feature_names = tuple(booster.feature_names or ()) predict = None if request.output_target != "log_loss": predict_type = "margin" if model_output in {None, "raw"} else "value" @@ -216,6 +370,7 @@ def predict(values: np.ndarray) -> np.ndarray: feature_map_digest=context.feature_map_digest, ) + shap = _require_shap() metadata = context.metadata or {} reference_data = metadata.get("reference_data") if reference_data is None: @@ -273,23 +428,28 @@ def explain_batch( ) if request.output_target == "log_loss": model_outputs = values.sum(axis=1) + base_values + elif getattr(explanation, "model_outputs", None) is not None: + model_outputs = _normalise_model_outputs( + explanation.model_outputs, + values, + ) elif prepared.predict is not None: model_outputs = _normalise_model_outputs(prepared.predict(batch), values) - if prepared.backend == "tree": - reconstructed = values.sum(axis=1) + base_values - if not np.allclose( - reconstructed, - model_outputs, - rtol=1e-3, - atol=1e-4, - equal_nan=False, - ): - raise ValueError( - "Tree SHAP additivity check failed for the declared " - f"output_target={request.output_target!r}" - ) else: model_outputs = None + if prepared.backend == "tree" and model_outputs is not None: + reconstructed = values.sum(axis=1) + base_values + if not np.allclose( + reconstructed, + model_outputs, + rtol=1e-3, + atol=1e-4, + equal_nan=False, + ): + raise ValueError( + "Tree SHAP additivity check failed for the declared " + f"output_target={request.output_target!r}" + ) feature_names = prepared.feature_names or tuple( f"feature_{index}" for index in range(values.shape[1]) ) @@ -297,11 +457,24 @@ def explain_batch( feature_names = tuple( f"feature_{index}" for index in range(values.shape[1]) ) + output_indexes = _selected_output_indexes( + request, + model_outputs=model_outputs, + row_count=values.shape[0], + output_count=values.shape[2], + ) + ranking_values = values + if request.output_selection == "predicted": + ranking_values = np.take_along_axis( + values, + output_indexes[:, None, :], + axis=2, + ) max_features = ( request.limits.top_k or request.limits.max_features or values.shape[1] ) if max_features < values.shape[1]: - feature_indexes = np.argsort(-np.abs(values).max(axis=2), axis=1)[ + feature_indexes = np.argsort(-np.abs(ranking_values).max(axis=2), axis=1)[ :, :max_features ] else: @@ -310,25 +483,31 @@ def explain_batch( rows: list[FeatureAttribution] = [] for row_index, input_id in enumerate(input_ids): for feature_index in feature_indexes[row_index]: - for output_index in range(values.shape[2]): + for output_index in output_indexes[row_index]: + feature_position = int(feature_index) + output_position = int(output_index) rows.append( FeatureAttribution( input_id=input_id, - output_id=f"output_{output_index}", - feature_id=str(feature_names[feature_index]), - feature_name=str(feature_names[feature_index]), + output_id=f"output_{output_position}", + feature_id=str(feature_names[feature_position]), + feature_name=str(feature_names[feature_position]), feature_view=request.feature_view, feature_value=( - _scalar(data_array[row_index, feature_index]) + _scalar(data_array[row_index, feature_position]) if request.result_policy.allow_sensitive_features else None ), contribution=float( - values[row_index, feature_index, output_index] + values[ + row_index, + feature_position, + output_position, + ] ), - base_value=float(base_values[row_index, output_index]), + base_value=float(base_values[row_index, output_position]), model_output=( - _scalar(model_outputs[row_index, output_index]) + _scalar(model_outputs[row_index, output_position]) if model_outputs is not None else None ), @@ -384,6 +563,51 @@ def _supported_objective(objective: str) -> bool: } +def _xgboost_booster_kind(booster: Any) -> str: + try: + config = json.loads(booster.save_config()) + kind = config["learner"]["gradient_booster"]["name"] + except (AttributeError, KeyError, TypeError, ValueError, json.JSONDecodeError): + raise ValueError( + "XGBoost booster type metadata is missing or invalid" + ) from None + if not isinstance(kind, str) or not kind: + raise ValueError("XGBoost booster type metadata is missing or invalid") + return kind + + +def _uses_native_tree_strategy( + context: ExplainableModelContext, + request: ExplainabilityRequest, +) -> bool: + return ( + context.flavor_id == "xgboost-native-v1" + and request.output_target in _NATIVE_OUTPUT_TARGETS + and request.reference is None + ) + + +def _selected_output_indexes( + request: ExplainabilityRequest, + *, + model_outputs: np.ndarray | None, + row_count: int, + output_count: int, +) -> np.ndarray: + if request.output_selection == "all": + return np.tile(np.arange(output_count), (row_count, 1)) + if model_outputs is None: + raise ValueError("output_selection='predicted' requires model outputs") + if model_outputs.shape != (row_count, output_count): + raise ValueError( + "model outputs do not match the attribution output shape for " + "output_selection='predicted'" + ) + if output_count == 1: + return np.zeros((row_count, 1), dtype=np.int64) + return np.argmax(model_outputs, axis=1)[:, None] + + def _tree_model_output(output_target: str) -> str | None: return { "model_output": None, diff --git a/tests/explainability/test_contracts.py b/tests/explainability/test_contracts.py index eb7320c..8fa8ad2 100644 --- a/tests/explainability/test_contracts.py +++ b/tests/explainability/test_contracts.py @@ -56,6 +56,18 @@ def test_request_model_role_is_resolved_from_bundle_descriptor_by_default() -> N assert request.model_role is None +def test_output_selection_defaults_to_all_and_is_serializable() -> None: + request = _request() + assert request.output_selection == "all" + assert request.model_dump(mode="json")["output_selection"] == "all" + + predicted = _request(output_selection="predicted") + assert predicted.output_selection == "predicted" + + with pytest.raises(ValidationError, match="output_selection"): + _request(output_selection="unsupported") + + def test_operation_store_uri_is_local_and_credential_free() -> None: request = _request(operation_store_uri="file:///tmp/tributo-operations") assert request.operation_store_uri == "file:///tmp/tributo-operations" @@ -271,4 +283,35 @@ def test_result_policy_is_explicit_and_serializable() -> None: def test_planner_preflights_known_explanation_byte_budget() -> None: request = _request(limits={"max_explanation_bytes": 100}) with pytest.raises(ValueError, match="estimated explanation output"): - ExplainabilityPlanner.preflight_limits(request, input_rows=10) + ExplainabilityPlanner.preflight_limits( + request, + input_rows=10, + output_count=2, + ) + + +def test_planner_uses_explicit_multiclass_output_bound_and_selection() -> None: + all_outputs = ExplainabilityPlanner.preflight_limits( + _request(), + input_rows=10, + output_count=10, + ) + predicted = ExplainabilityPlanner.preflight_limits( + _request(output_selection="predicted"), + input_rows=10, + output_count=10, + ) + + assert all_outputs["estimated_output_count"] == 10 + assert all_outputs["estimated_output_rows"] == 200 + assert predicted["estimated_output_count"] == 1 + assert predicted["estimated_output_rows"] == 20 + + +def test_planner_rejects_non_positive_output_count() -> None: + with pytest.raises(ValueError, match="output_count"): + ExplainabilityPlanner.preflight_limits( + _request(), + input_rows=10, + output_count=0, + ) diff --git a/tests/explainability/test_executor.py b/tests/explainability/test_executor.py index 9a177a7..6cbe245 100644 --- a/tests/explainability/test_executor.py +++ b/tests/explainability/test_executor.py @@ -20,9 +20,12 @@ from tributo.explainability.executor import ( _attempt_result_uri, _build_onnx_inputs, + _explanation_output_count_upper_bound, _LeaseHeartbeat, _load_reference, + _make_receipt, _manifest_role_digest, + _operation_idempotency_key, _operation_store_for_request, _resolve_xgboost_feature_names, _schema_signature, @@ -158,6 +161,63 @@ def test_tree_descriptor_resolves_explainability_role_when_request_omits_role() _validate_request_against_descriptor(manifest, request) +@pytest.mark.parametrize( + ("task_type", "field_name", "shape", "expected"), + [ + ("classification", "probabilities", ("batch", 10), 10), + ("classification", "probabilities", ("batch", 2), 2), + ("regression", "prediction", ("batch", 1), 1), + ], +) +def test_xgboost_output_bound_comes_from_typed_manifest_signature( + task_type: str, + field_name: str, + shape: tuple[str | int, ...], + expected: int, +) -> None: + request = ExplainabilityRequest( + bundle_uri="/models/bundle", + input=IngestionRequest( + source=ParquetSourceConfig(path="/data/input.parquet"), engine="ray" + ), + backend="tree", + result_uri="/data/explanations", + request_id="request-output-bound", + ) + artifact = SimpleNamespace(name="native", flavor_id="xgboost-native-v1") + manifest = SimpleNamespace( + roles={"explainability_model": "native"}, + artifacts=(artifact,), + source_info=SimpleNamespace(task_type=task_type), + output_signature=SimpleNamespace( + output_fields=(SimpleNamespace(name=field_name, shape=shape),) + ), + ) + + assert _explanation_output_count_upper_bound(manifest, request) == expected + + +def test_xgboost_output_bound_requires_a_fixed_typed_signature() -> None: + request = ExplainabilityRequest( + bundle_uri="/models/bundle", + input=IngestionRequest( + source=ParquetSourceConfig(path="/data/input.parquet"), engine="ray" + ), + backend="tree", + result_uri="/data/explanations", + request_id="request-missing-output-bound", + ) + manifest = SimpleNamespace( + roles={"explainability_model": "native"}, + artifacts=(SimpleNamespace(name="native", flavor_id="xgboost-native-v1"),), + source_info=SimpleNamespace(task_type="classification"), + output_signature=SimpleNamespace(output_fields=()), + ) + + with pytest.raises(ValueError, match="typed probability or prediction"): + _explanation_output_count_upper_bound(manifest, request) + + def test_descriptorless_bundle_is_rejected_before_worker_loading() -> None: request = ExplainabilityRequest( bundle_uri="/models/bundle", @@ -258,6 +318,11 @@ def open(self, _selection): "_validate_request_against_descriptor", lambda _manifest, _request: None, ) + monkeypatch.setattr( + executor_module, + "_explanation_output_count_upper_bound", + lambda _manifest, _request: 1, + ) monkeypatch.setattr( executor_module, "_selected_backend", @@ -308,6 +373,61 @@ def fake_make_receipt(**kwargs): assert record.result_uri != request.result_uri +def test_receipt_and_idempotency_record_output_selection() -> None: + request = ExplainabilityRequest( + bundle_uri="/models/bundle", + input=IngestionRequest( + source=ParquetSourceConfig(path="/data/input.parquet"), engine="ray" + ), + backend="tree", + output_selection="predicted", + result_uri="/data/explanations", + request_id="request-output-selection", + ) + artifact = SimpleNamespace( + name="native", + flavor_id="xgboost-native-v1", + tree_digest="b" * 64, + files=(), + ) + manifest = SimpleNamespace( + bundle_id="bundle-output-selection", + roles={"explainability_model": "native"}, + artifacts=(artifact,), + explainability=None, + ) + receipt = _make_receipt( + manifest=manifest, + request=request, + operation_id="operation-output-selection", + bundle_digest="a" * 64, + selected_backend="tree", + exactness="exact", + input_rows=1, + explanation_rows=2, + result_digest="c" * 64, + result_bytes=128, + result_uri="/data/explanations/attempts/lease", + status="succeeded", + reference_provider=SimpleNamespace(), + ) + assert receipt.output_selection == "predicted" + + all_key = _operation_idempotency_key( + manifest, + request.model_copy(update={"output_selection": "all"}), + bundle_digest="a" * 64, + reference_provider=SimpleNamespace(), + ) + predicted_key = _operation_idempotency_key( + manifest, + request, + bundle_digest="a" * 64, + reference_provider=SimpleNamespace(), + ) + assert all_key != predicted_key + + def test_xgboost_feature_order_is_checked_without_sidecar() -> None: class FakeBooster: feature_names = ["feature_a", "feature_b"] diff --git a/tests/explainability/test_shap.py b/tests/explainability/test_shap.py index 3fb064e..78525e8 100644 --- a/tests/explainability/test_shap.py +++ b/tests/explainability/test_shap.py @@ -7,8 +7,9 @@ import numpy as np import pytest +from tributo.explainability import shap as shap_module from tributo.explainability.protocols import ExplainableModelContext, PreparedExplainer -from tributo.explainability.shap import ShapAdapter +from tributo.explainability.shap import ShapAdapter, _NativeTreeExplainer from .test_contracts import _request @@ -41,6 +42,18 @@ def _base_value(label: float) -> float: return 0.5 if label == 0 else 0.6 +class _MultiOutputExplanation: + values = np.asarray( + [ + [[10.0, 0.1], [0.2, 5.0]], + [[4.0, 0.1], [0.2, 3.0]], + ] + ) + data = np.asarray([[10.0, 20.0], [30.0, 40.0]]) + base_values = np.asarray([[-10.0, 0.0], [0.0, -3.0]]) + model_outputs = values.sum(axis=1) + base_values + + def test_shap_long_rows_use_top_k_and_preserve_provenance() -> None: request = _request(limits={"top_k": 1}) prepared = PreparedExplainer( @@ -64,6 +77,51 @@ def test_shap_long_rows_use_top_k_and_preserve_provenance() -> None: assert all(row.feature_value is None for row in rows) +def test_predicted_selection_preserves_class_id_and_ranks_selected_output() -> None: + request = _request(output_selection="predicted", limits={"top_k": 1}) + prepared = PreparedExplainer( + backend="tree", + exactness="exact", + explain=lambda batch, **_: _MultiOutputExplanation(), + 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, + ) + + assert [(row.feature_name, row.output_id) for row in rows] == [ + ("feature_b", "output_1"), + ("feature_a", "output_0"), + ] + + +def test_binary_predicted_selection_is_the_single_output() -> None: + request = _request(output_selection="predicted") + prepared = PreparedExplainer( + backend="tree", + exactness="exact", + explain=lambda batch, **_: _Explanation(), + feature_names=("feature_a", "feature_b"), + predict=lambda batch: np.asarray([1.6, 2.7]), + ) + + 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, + ) + + assert len(rows) == 4 + assert {row.output_id for row in rows} == {"output_0"} + + def test_tree_log_loss_requires_labels_at_adapter_boundary() -> None: request = _request( backend="tree", @@ -220,6 +278,233 @@ def test_tree_support_rejects_unknown_output_target() -> None: assert "output_target" in decision.reason +def test_tree_support_rejects_predicted_selection_outside_native_classification() -> ( + None +): + regression = ExplainableModelContext( + bundle_uri="/models/bundle", + model_role="inference", + artifact_name="native", + artifact_format="ubj", + flavor_id="xgboost-native-v1", + artifact_path=None, + objective="reg:squarederror", + ) + regression_decision = ShapAdapter.supports( + regression, + _request(backend="tree", output_selection="predicted"), + ) + assert regression_decision.supported is False + assert "classification objective" in regression_decision.reason + + classification = replace(regression, objective="binary:logistic") + probability_decision = ShapAdapter.supports( + classification, + _request( + backend="tree", + output_target="probability", + output_selection="predicted", + reference={"uri": "/reference.npy"}, + ), + ) + assert probability_decision.supported is False + assert "native XGBoost raw" in probability_decision.reason + + +def test_native_prepare_does_not_load_shap(monkeypatch) -> None: + pytest.importorskip("xgboost") + + class FakeBooster: + feature_names = ["feature_a", "feature_b"] + + @staticmethod + def save_config(): + return '{"learner":{"gradient_booster":{"name":"gbtree"}}}' + + context = ExplainableModelContext( + bundle_uri="/models/bundle", + model_role="inference", + artifact_name="native", + artifact_format="ubj", + flavor_id="xgboost-native-v1", + artifact_path=None, + model_object=FakeBooster(), + feature_names=("feature_a", "feature_b"), + objective="binary:logistic", + ) + + def fail_if_loaded(): + raise AssertionError("native XGBoost TreeSHAP must not load SHAP") + + monkeypatch.setattr(shap_module, "_require_shap", fail_if_loaded) + prepared = ShapAdapter().prepare(context, _request(backend="tree")) + assert isinstance(prepared.explain, _NativeTreeExplainer) + + +def test_native_prepare_rejects_non_tree_booster() -> None: + pytest.importorskip("xgboost") + + class FakeBooster: + feature_names = ["feature_a", "feature_b"] + + @staticmethod + def save_config(): + return '{"learner":{"gradient_booster":{"name":"gblinear"}}}' + + context = ExplainableModelContext( + bundle_uri="/models/bundle", + model_role="inference", + artifact_name="native", + artifact_format="ubj", + flavor_id="xgboost-native-v1", + artifact_path=None, + model_object=FakeBooster(), + feature_names=("feature_a", "feature_b"), + objective="binary:logistic", + ) + + with pytest.raises(ValueError, match="gbtree or dart"): + ShapAdapter().prepare(context, _request(backend="tree")) + + +def test_native_tree_shap_rejects_non_strict_contribution_shape() -> None: + pytest.importorskip("xgboost") + + class FakeBooster: + feature_types = None + + @staticmethod + def predict(matrix, **kwargs): + rows = matrix.num_row() + if kwargs.get("pred_contribs"): + return np.zeros((rows, 3), dtype=np.float32) + return np.zeros((rows, 1), dtype=np.float32) + + explainer = _NativeTreeExplainer( + FakeBooster(), + feature_names=("feature_a", "feature_b"), + objective="binary:logistic", + ) + with pytest.raises(ValueError, match="strict shape contract"): + explainer(np.asarray([[0.0, 1.0]], dtype=np.float32)) + + +def test_real_xgboost_native_tree_shap_supports_regression() -> None: + xgboost = pytest.importorskip("xgboost") + X = np.asarray([[0.0, 0.0], [0.0, 1.0], [1.0, 0.0], [1.0, 1.0]]) + y = np.asarray([0.0, 1.0, 1.0, 2.0]) + matrix = xgboost.DMatrix( + X, + label=y, + feature_names=["feature_a", "feature_b"], + ) + booster = xgboost.train( + {"objective": "reg:squarederror", "max_depth": 2, "eta": 0.5}, + matrix, + num_boost_round=4, + ) + context = ExplainableModelContext( + bundle_uri="/models/bundle", + model_role="native", + artifact_name="native", + artifact_format="ubj", + flavor_id="xgboost-native-v1", + artifact_path=None, + model_object=booster, + feature_names=("feature_a", "feature_b"), + objective="reg:squarederror", + ) + request = _request(backend="tree", output_target="raw_margin") + + rows = ShapAdapter().explain_batch( + ShapAdapter().prepare(context, request), + X, + input_ids=("0", "1", "2", "3"), + model_digest="a" * 64, + request=request, + ) + + assert len(rows) == len(X) * X.shape[1] + assert {row.output_id for row in rows} == {"output_0"} + expected = booster.predict(matrix, output_margin=True, strict_shape=True) + for row_index in range(len(X)): + selected = rows[row_index * X.shape[1] : (row_index + 1) * X.shape[1]] + reconstructed = sum(row.contribution for row in selected) + float( + selected[0].base_value + ) + assert reconstructed == pytest.approx(float(expected[row_index, 0])) + + +def test_real_xgboost_native_tree_shap_selects_multiclass_output() -> None: + xgboost = pytest.importorskip("xgboost") + X = np.asarray( + [ + [0.0, 0.0], + [0.0, 1.0], + [1.0, 0.0], + [1.0, 1.0], + [2.0, 0.0], + [2.0, 1.0], + ], + dtype=np.float32, + ) + y = np.asarray([0, 0, 1, 1, 2, 2], dtype=np.float32) + matrix = xgboost.DMatrix( + X, + label=y, + feature_names=["feature_a", "feature_b"], + ) + booster = xgboost.train( + { + "objective": "multi:softprob", + "num_class": 3, + "max_depth": 2, + "eta": 0.5, + }, + matrix, + num_boost_round=6, + ) + context = ExplainableModelContext( + bundle_uri="/models/bundle", + model_role="native", + artifact_name="native", + artifact_format="ubj", + flavor_id="xgboost-native-v1", + artifact_path=None, + model_object=booster, + feature_names=("feature_a", "feature_b"), + objective="multi:softprob", + ) + all_request = _request(backend="tree") + predicted_request = _request( + backend="tree", + output_selection="predicted", + ) + adapter = ShapAdapter() + all_rows = adapter.explain_batch( + adapter.prepare(context, all_request), + X, + input_ids=tuple(str(index) for index in range(len(X))), + model_digest="a" * 64, + request=all_request, + ) + predicted_rows = adapter.explain_batch( + adapter.prepare(context, predicted_request), + X, + input_ids=tuple(str(index) for index in range(len(X))), + model_digest="a" * 64, + request=predicted_request, + ) + + assert len(all_rows) == len(X) * X.shape[1] * 3 + assert len(predicted_rows) == len(X) * X.shape[1] + margins = booster.predict(matrix, output_margin=True, strict_shape=True) + expected_outputs = np.argmax(margins, axis=1) + for row_index, output_index in enumerate(expected_outputs): + selected = predicted_rows[row_index * X.shape[1] : (row_index + 1) * X.shape[1]] + assert {row.output_id for row in selected} == {f"output_{output_index}"} + + def test_real_xgboost_tree_shap_checks_raw_and_probability_outputs() -> None: xgboost = pytest.importorskip("xgboost") pytest.importorskip("shap") diff --git a/tests/integration/test_explainability_ray_jobs.py b/tests/integration/test_explainability_ray_jobs.py index c9151b8..ef5fb29 100644 --- a/tests/integration/test_explainability_ray_jobs.py +++ b/tests/integration/test_explainability_ray_jobs.py @@ -28,7 +28,6 @@ from tributo.exporting.manifest import ( ExportManifestV2, ManifestExecution, - ManifestSignature, ManifestSourceInfo, ) from tributo.exporting.models import ( @@ -208,7 +207,7 @@ def test_onnx_model_agnostic_shap_ray_job_writes_long_parquet_and_receipt( def test_xgboost_ubj_tree_shap_ray_job_writes_exact_attributions( objective: str, ) -> None: - """Run the real XGBoost UBJ + TreeExplainer path in a Ray actor.""" + """Run the native XGBoost TreeSHAP path in a Ray actor.""" import xgboost shared_root = Path( @@ -230,13 +229,20 @@ def test_xgboost_ubj_tree_shap_ray_job_writes_exact_attributions( "max_depth": 2, "eta": 0.5, } + class_count = 2 if objective == "multi:softprob": - training_params["num_class"] = 3 + class_count = 3 + training_params["num_class"] = class_count booster = xgboost.train( training_params, xgboost.DMatrix(X, label=y, feature_names=["feature_a", "feature_b"]), num_boost_round=4, ) + expected_margins = booster.predict( + xgboost.DMatrix(X, feature_names=["feature_a", "feature_b"]), + output_margin=True, + strict_shape=True, + ) model_bytes = bytes(booster.save_raw(raw_format="ubj")) artifact_dir = root / "bundle" / "artifacts" / "native" artifact_dir.mkdir(parents=True) @@ -262,14 +268,45 @@ def test_xgboost_ubj_tree_shap_ray_job_writes_exact_attributions( backend="tree", model_role="explainability_model", ) + checkpoint_contract = ExportCheckpointV1( + trainer_type="xgboost", + architecture_id="xgboost", + input_schema=( + CheckpointField( + name="float_input", + dtype="float32", + shape=("batch", X.shape[1]), + ), + ), + output_schema=( + CheckpointField(name="label", dtype="int64", shape=("batch",)), + CheckpointField( + name="probabilities", + dtype="float32", + shape=("batch", class_count), + ), + ), + task_type="classification", + framework="xgboost", + framework_version=xgboost.__version__, + preprocessing={"type": "none"}, + checkpoint_format_version=1, + ) + input_signature, output_signature = checkpoint_contract.to_manifest_signatures() manifest = ExportManifestV2( bundle_id="bundle-tree-it", status="succeeded", canonical_uri=str(root / "bundle"), tributo_version="1.0.0", - source_info=ManifestSourceInfo(source_kind="xgboost_result"), - input_signature=ManifestSignature(), - output_signature=ManifestSignature(), + source_info=ManifestSourceInfo( + source_kind="xgboost_result", + framework=checkpoint_contract.framework, + framework_version=checkpoint_contract.framework_version, + architecture_id=checkpoint_contract.architecture_id, + task_type=checkpoint_contract.task_type, + ), + input_signature=input_signature, + output_signature=output_signature, artifacts=(artifact,), roles={"explainability_model": "native"}, execution=ManifestExecution(execution_id="exec-tree-it"), @@ -290,6 +327,7 @@ def test_xgboost_ubj_tree_shap_ray_job_writes_exact_attributions( input_dir / "part-0.parquet", ) result_dir = root / "tree-explanations" + output_selection = "predicted" if objective == "multi:softprob" else "all" request = ExplainabilityRequest( bundle_uri=str(root / "bundle"), input=IngestionRequest( @@ -299,6 +337,7 @@ def test_xgboost_ubj_tree_shap_ray_job_writes_exact_attributions( feature_columns=("feature_a", "feature_b"), input_id_column="entity_id", backend="tree", + output_selection=output_selection, result_uri=str(result_dir), operation_store_uri=str(root / "operations"), request_id=f"tree-request-{uuid.uuid4().hex}", @@ -327,7 +366,8 @@ def test_xgboost_ubj_tree_shap_ray_job_writes_exact_attributions( assert receipt["status"] == "succeeded" assert receipt["backend"] == "tree" assert receipt["exactness"] == "exact" - assert receipt["explanation_rows"] > 0 + assert receipt["output_selection"] == output_selection + assert receipt["explanation_rows"] == len(X) * X.shape[1] result_files = sorted(result_path.glob("*.parquet")) result_table = pa.concat_tables([pq.read_table(path) for path in result_files]) assert result_table.num_rows == receipt["explanation_rows"] @@ -335,6 +375,27 @@ def test_xgboost_ubj_tree_shap_ray_job_writes_exact_attributions( "feature_a", "feature_b", } + result_rows = result_table.to_pylist() + expected_output_indexes = ( + np.argmax(expected_margins, axis=1) + if objective == "multi:softprob" + else np.zeros(len(X), dtype=np.int64) + ) + for row_index, (input_id, output_index) in enumerate( + zip((201, 202, 203, 204), expected_output_indexes, strict=True) + ): + selected = [row for row in result_rows if row["input_id"] == str(input_id)] + assert len(selected) == X.shape[1] + assert {row["output_id"] for row in selected} == {f"output_{output_index}"} + assert len({row["base_value"] for row in selected}) == 1 + assert len({row["model_output"] for row in selected}) == 1 + reconstructed = ( + sum(row["contribution"] for row in selected) + selected[0]["base_value"] + ) + assert reconstructed == pytest.approx(selected[0]["model_output"]) + assert selected[0]["model_output"] == pytest.approx( + expected_margins[row_index, output_index] + ) _assert_persisted_succeeded_operation(root, request.request_id) finally: import shutil