Interpretability
BaseModel provides model interpretation via gradient-based attribution — either Integrated Gradients (deterministic, default) or GradientSHAP (stochastic, faster) — generating feature-level and event-level attribution analysis. This helps explain why the model makes specific predictions for each entity. The same interpret() entry point can also render a client-ready SHAP-library-style report (beeswarm, bar, heatmap, waterfall, force plots) alongside the standard outputs.
Functions
from monad.interpretability import (
interpret,
interpret_entity,
attributions_to_shap_explanation,
save_shap_report,
)
from monad.interpretability import TreemapGenerator
interpret()
Generate aggregate interpretability analysis across many entities. Produces feature importance plots and JSON summaries.
from pathlib import Path
from datetime import datetime, timezone
from monad.interpretability import interpret
interpret(
predictions_path=Path("./predictions.tsv"),
output_path=Path("./interpretations"),
checkpoint_path=Path("./my_model"),
device="cuda",
)
Parameters
| Parameter | Type | Default | Description |
|---|---|---|---|
predictions_path | Path | required | Path to TSV file with model predictions. |
output_path | Path | required | Directory to store interpretation results. |
checkpoint_path | Path | required | Path to the scenario model checkpoint. |
device | str | required | Device for computation: "cpu" or "cuda". |
prediction_date | datetime | datetime.now(tz=timezone.utc) | Date for which predictions are interpreted. Uses UTC timezone. |
limit_batches | int \| None | None | Number of predict batches to attribute. None = all batches. Also caps the batches the reference is drawn from. Together with batch_size, it caps the sample at limit_batches × batch_size observations. |
target_index | int \| None | None | Output index to attribute. Required for classification models (binary, multiclass, multilabel); use 0 for binary classification. Optional for single-target regression; required for multi-target regression (num_targets > 1) to select the target being explained. Ignored for recommendation, where the index of recommended_value is used. |
classification_resample | bool | False | Resample data for balanced classes. Only for classification models. |
recommended_value | str \| None | None | Item ID to interpret. Only for recommendation models. |
group_size | int | 500 | Maximum samples per group. For classification: per class. For recommendation: total observations. |
method | Literal["integrated_gradients", "gradient_shap"] | "integrated_gradients" | Attribution method. "integrated_gradients" runs a deterministic path integral; "gradient_shap" averages gradients across n_samples Gaussian-perturbed baselines (stochastic, typically faster, denser attributions). Both operate on the input-space Cleora/EMDE float sketches. See Attribution Methods. |
save_shap_plots | bool | False | If True, additionally render a SHAP-library-style report (beeswarm / bar / heatmap / waterfall / force + static shap_report.html index) under <output_path>/shap/. Requires the shap optional extra. See SHAP-Style Report. |
top_n | int \| None | None | If set, top_features.json and top_features.png keep only the top-N features, ranked across all data sources. None = all features. |
batch_size | int \| None | None | Rows per batch of the predict walk. None = the batch size the checkpoint was fitted with. While limit_batches is set, changing it also changes how many observations are attributed. See Controlling Memory. |
n_steps | int | 50 | Integrated Gradients only. Number of steps the path integral is approximated over. Fewer steps give a coarser integral and change the attributions. |
steps_per_pass | int \| None | None | Integrated Gradients only. How many of the n_steps steps go through the model together. None = all of them at once. Peak memory is about steps_per_pass × batch_size rows; the attributions do not change. Must not exceed n_steps. |
n_samples | int | 25 | GradientSHAP only. Number of perturbed samples the gradients are averaged over. Fewer samples give noisier attributions; this also sets runtime and peak memory. |
n_baselines | int | 100 | GradientSHAP only. Number of real observations kept as the reference distribution. More rows give a steadier result and use more device memory. |
seed | int | 42 | Seed for the GradientSHAP reference draw and for its random sampling. Integrated Gradients does not use it. |
Parameters are validated when interpret() is called, before the checkpoint is loaded. A non-positive group_size, top_n, limit_batches, batch_size, n_steps, steps_per_pass, n_samples or n_baselines, a negative target_index, or a steps_per_pass larger than n_steps raises a ValueError (a pydantic ValidationError).
Controlling Memory
Attribution expands each batch internally, so at the batch size the checkpoint was fitted with, a run can need more memory than one device holds. If a run does not fit:
steps_per_pass(Integrated Gradients): try this first. Peak memory is aboutsteps_per_pass × batch_sizerows, and neither the attributions nor the sample change.batch_size: lower the number of rows per batch. Withlimit_batchesset, this also shrinks the sample, because the run reads at mostlimit_batches × batch_sizeobservations.limit_batches: read fewer batches, for both the attributed observations and the reference.
For GradientSHAP, steps_per_pass has no effect; the batch width bounds each pass, and n_samples and n_baselines also set its peak memory.
interpret(
predictions_path=Path("./predictions.tsv"),
output_path=Path("./interpretations"),
checkpoint_path=Path("./my_model"),
device="cuda",
batch_size=32,
steps_per_pass=5,
)
Output Files
The output_path directory will contain a nested structure:
output_path/
├── run.json # What the run was measured against (written last)
├── source_importance.json # Attribution scores per data source
├── source_importance.png # Bar chart of data source importance
├── top_features.json # Features ranked across all data sources (see top_n)
├── top_features.png # Bar chart of the top features
├── {data_source}/ # One directory per data source
│ ├── feature_importance.json # Attribution scores per feature
│ ├── feature_importance.png # Bar chart of feature importance
│ ├── Events frequency/ # Only when event counts span several time buckets
│ │ └── buckets_importance.png # Per-bucket breakdown of event counts
│ └── {feature}/ # One directory per feature
│ ├── values_importance.json # Attribution scores per feature value
│ ├── values_collision.json # Categorical and recency sketches only
│ ├── values_highest_importance.png # Top positive attributions
│ └── values_lowest_importance.png # Top negative attributions
└── shap/ # Only when save_shap_plots=True
├── shap_beeswarm.png # Per-feature signed-distribution beeswarm
├── shap_bar_global.png # Global mean-absolute attribution bar
├── shap_heatmap.png # Per-entity × per-feature heatmap (n_entities ≥ 10)
├── shap_waterfall_top{i}.png # Top-N entity waterfall plots
├── shap_force_top{i}.html # Top-N entity interactive force plots
└── shap_report.html # Static index linking all artifacts
run.json is written after every other artifact, so a run that fails midway does not write one. It contains:
| Key | Description |
|---|---|
method | Attribution method used. |
reference | Identifier of the reference the attributions are measured against. |
observations | Number of observations attributed. |
base_value | The model's average output over the reference rows. |
prediction_date | Prediction date, in ISO 8601 format. |
target | The target_index passed to the run, or null. |
n_steps | The run's n_steps setting (used by Integrated Gradients only). |
comparability | Fixed note on how runs can be compared (see below). |
Every importance chart (the PNG files outside shap/) carries a subtitle with the same method, reference and observation count, and labels each bar with its value and its net share: the fraction of the bar's magnitude that points in one direction. A net share near 0% means positive and negative contributions cancel out. Magnitudes are comparable between methods and between runs. Directions (sign and bar colour) are relative to the reference, so compare them only between runs with the same reference: a different reference proves two runs' directions are not comparable, and a matching one is strong evidence, not proof, that they are.
values_collision.json is written for features encoded with a categorical or recency sketch. A sketch stores values at positions shared out of a fixed pool, so distinct values can land in the same position, and the value breakdown in values_importance.json then credits each of them with what belongs to the group. The file reports how much of the breakdown is affected: value_count, value_positions, occupied_positions, shared_positions, colliding_value_count, colliding_value_fraction (the share of values that share a position with another value), shared_position_fraction, and values_per_occupied_position (the average number of values per occupied position; 1.0 means every value has its own). The importances themselves are unchanged. Decimal features get no file, because their values share quantile buckets by design.
Attribution Methods
interpret() supports two gradient-based attribution methods via the method parameter. Both operate on the same input — the pre-computed Cleora/EMDE sketches monad consumes — and produce attributions in the same layout, so the downstream output structure is identical.
| Method | method= value | Determinism | Typical runtime | Output density | Recommended for |
|---|---|---|---|---|---|
| Integrated Gradients | "integrated_gradients" | Deterministic (path integral) | Baseline | Sparser | Reproducible single runs, regulated reporting |
| GradientSHAP | "gradient_shap" | Stochastic (seeded; reproducible per-call) | ~2× faster on representative workloads | Denser | Production-scale attribution, SHAP-style downstream plots |
Switch methods with one keyword:
from monad.interpretability import interpret
interpret(
predictions_path=Path("./predictions.tsv"),
output_path=Path("./interpretations"),
checkpoint_path=Path("./my_model"),
device="cuda",
method="gradient_shap",
)
Tune GradientSHAP with n_samples, n_baselines and seed, and Integrated Gradients with n_steps and steps_per_pass, all passed directly to interpret() (see Parameters). The standard deviation of the Gaussian noise GradientSHAP adds is fixed at 0.15.
Reference Point
Both methods measure attributions against a reference drawn from the predict set itself:
- Integrated Gradients integrates from the average of the real observations the run reads (the per-position mean of the model inputs). The attributions of an observation add up, approximately, to its prediction minus the prediction at that average, so they describe how the observation differs from a typical one. In 1.13 and earlier, the reference was an all-zeros input, so Integrated Gradients attributions and
base_valuediffer from runs made with those releases. - GradientSHAP samples up to
n_baselinesreal observations at random, seeded byseed, as its background distribution.
The reference is drawn from the whole predict population, not only from the entities being explained, and from no more batches than limit_batches allows. With limit_batches set, the GradientSHAP background is sampled from those batches only.
The reference used is named in run.json and in every chart subtitle, for example typical_observation(rows=1,fingerprint=…) for Integrated Gradients or sampled_distribution(rows=100,seed=42,fingerprint=…) for GradientSHAP.
Completeness warning
With Integrated Gradients, the first batch is checked for completeness: each observation's attributions should add up to the difference between its prediction and base_value. If any observation's attributions miss that difference by more than 5% of it, a warning is logged (Attributions miss the prediction by …), meaning the attributions' relative sizes carry more meaning than their values. Raise n_steps to tighten the path integral. GradientSHAP is complete only in expectation, so it is not checked.
SHAP-Style Report
When save_shap_plots=True, interpret() writes a parallel shap/ subdirectory with a SHAP-library-style report (beeswarm, global bar, heatmap, per-entity waterfall and interactive force plots), plus a static shap_report.html index that links everything. The report is designed for client hand-off and renders without a Jupyter runtime.
from monad.interpretability import interpret
interpret(
predictions_path=Path("./predictions.tsv"),
output_path=Path("./interpretations"),
checkpoint_path=Path("./my_model"),
device="cuda",
method="gradient_shap",
save_shap_plots=True,
)
Optional dependency
The SHAP report requires the shap library, shipped in the interpretability extra. Install with poetry install -E interpretability or pip install '.[interpretability]'. Without it, interpret() raises ImportError only when save_shap_plots=True; the default attribution flow remains unaffected.
For rendering the report from already-computed attributions (e.g., a custom batch pipeline), use the two helpers below directly.
attributions_to_shap_explanation()
Convert monad per-entity attribution tensors into a shap.Explanation object that can be passed to any shap.plots.* function.
from monad.interpretability import attributions_to_shap_explanation
explanation = attributions_to_shap_explanation(
attributions=attributions,
feature_slicer=feature_slicer,
base_value=base_value,
)
| Parameter | Type | Default | Description |
|---|---|---|---|
attributions | torch.Tensor | required | Attribution matrix shaped (n_entities, features), such as AttributionResults.attributions. |
feature_slicer | FeatureSlicer | required | The trained model's feature layout, the one the attributions were computed over. |
base_value | float | required | Scalar E[f(reference)] — the model's average output over the reference rows, such as AttributionResults.base_value. |
feature_values | torch.Tensor \| None | None | Raw inputs with the same shape and row order as attributions, such as AttributionResults.feature_values. When given, the beeswarm colours points by feature value; otherwise points are uncoloured. |
output_name | str \| None | None | Optional label stamped onto explanation.output_names (e.g. "target_0" for classification). |
Returns a shap.Explanation with values of shape (n_entities, n_features), signed-summed per feature slice, and a stable feature-name order keyed by "<data_source>.<feature>" to avoid cross-source name collisions.
Raises ValueError if attributions contains no entities, the slicer has no slices, or feature_values produces a different feature layout than attributions; ImportError if shap is not installed.
save_shap_report()
Render the full SHAP-style report from a shap.Explanation.
from monad.interpretability import save_shap_report
save_shap_report(
explanation=explanation,
output_path=Path("./interpretations"),
top_n_waterfalls=5,
)
| Parameter | Type | Default | Description |
|---|---|---|---|
explanation | shap.Explanation | required | A SHAP Explanation, typically from attributions_to_shap_explanation(). |
output_path | Path | required | Parent directory; a shap/ subdirectory is created under it. |
top_n_waterfalls | int | 5 | How many highest-impact entities to render waterfall and force plots for (capped at n_entities). Kept small by default because each force HTML embeds ~2 MB of shap.js. |
Returns the Path to the created shap/ directory. Writes shap_beeswarm.png, shap_bar_global.png, shap_heatmap.png (only when n_entities >= 10), shap_waterfall_top{i}.png, shap_force_top{i}.html, and shap_report.html. Individual plot failures are logged and skipped rather than aborting the whole report.
Raises ImportError if shap or matplotlib is not installed.
interpret_entity()
Explain a single entity's prediction at the event level. Produces a JSON file with per-event, per-feature attributions.
from pathlib import Path
from datetime import datetime, timezone
from monad.interpretability import interpret_entity
interpret_entity(
output_path=Path("./interpretations/customer_123.json"),
checkpoint_path=Path("./my_model"),
predictions_path=Path("./predictions.tsv"),
main_entity_id="customer_123",
device="cuda",
prediction_date=datetime(2024, 6, 1, tzinfo=timezone.utc),
)
Parameters
| Parameter | Type | Default | Description |
|---|---|---|---|
output_path | Path | required | Path for the output JSON file. |
checkpoint_path | Path | required | Path to the scenario model checkpoint. |
predictions_path | Path | required | Path to TSV file with model predictions. |
main_entity_id | str | required | Entity ID to interpret. For some databases (e.g., Snowflake), the value may need escaping. |
device | str | required | Device: "cpu" or "cuda". |
prediction_date | datetime | datetime.now(tz=timezone.utc) | Date for which to interpret. Uses UTC timezone. |
target_index | int \| None | None | Output index to attribute. Required for classification models (binary, multiclass, multilabel); use 0 for binary classification. Optional for single-target regression; required for multi-target regression (num_targets > 1) to select the target being explained. |
recommended_value | str \| None | None | Item ID to interpret (recommendation only). |
method | Literal["integrated_gradients", "gradient_shap"] | "integrated_gradients" | Attribution method. See Attribution Methods. |
batch_size | int \| None | None | Rows per batch of the predict walk. None = the batch size the checkpoint was fitted with. |
limit_batches | int \| None | None | Maximum number of predict batches read, both for the explained entity and for the reference. None = all batches. The explained entity fits in one batch, so in practice this caps the population the reference is drawn from, which keeps a single-entity run cheap on large datasets. |
n_steps | int | 50 | Integrated Gradients only. Same as in interpret(). |
steps_per_pass | int \| None | None | Integrated Gradients only. Same as in interpret(). |
n_samples | int | 25 | GradientSHAP only. Same as in interpret(). |
n_baselines | int | 100 | GradientSHAP only. Same as in interpret(). |
seed | int | 42 | Same as in interpret(). |
Parameters are validated when interpret_entity() is called, with the same rules as for interpret(). The reference is drawn from the whole predict population, so the entity is compared with a typical entity, not with itself.
Output Format
The output JSON contains event-level attributions:
{
"transactions": [
{
"timestamp": "2020-01-25T00:00:00",
"modality_attributions": [
{
"data_source_name": "transactions",
"name": "article_id",
"value": "0854796002",
"attribution": -0.124
},
{
"data_source_name": "transactions",
"name": "price",
"value": "0.017",
"attribution": -0.009
}
]
}
]
}
Next to the JSON file, interpret_entity() writes a run manifest named <output file stem>_<main_entity_id>_run.json (for example customer_123_customer_123_run.json), with the same keys as run.json.
Batch Processing Multiple Entities
from pathlib import Path
from datetime import datetime, timezone
from monad.interpretability import interpret_entity
entities_to_explain = ["cust_001", "cust_002", "cust_003"]
for entity_id in entities_to_explain:
interpret_entity(
output_path=Path(f"./interpretations/{entity_id}.json"),
checkpoint_path=Path("./my_model"),
predictions_path=Path("./predictions.tsv"),
main_entity_id=entity_id,
device="cuda",
prediction_date=datetime(2024, 6, 1, tzinfo=timezone.utc),
)
TreemapGenerator
Generate treemap visualizations from interpretability attribution data. Treemaps provide a hierarchical view of feature importance across data sources and features.
Constructor
| Parameter | Type | Default | Description |
|---|---|---|---|
interpretability_files_path | Path \| None | None | Path to directory with interpretation output (from interpret()). |
hierarchy | TreemapHierarchy \| None | None | Custom hierarchy definition for treemap levels. |
Note
Exactly one of interpretability_files_path or hierarchy must be provided.
plot_treemap()
Generate and save a treemap chart as an HTML file.
generator.plot_treemap(
output_file_path=Path("./treemap.html"),
n_largest_per_feature=1500,
max_depth=3,
)
| Parameter | Type | Default | Description |
|---|---|---|---|
output_file_path | Path | required | Path for the output HTML file. |
n_largest_per_feature | int \| None | 1500 | Maximum number of values per feature to include. |
n_largest | int \| None | None | Maximum total number of values to include. |
exclude_positive_attributions | bool | False | Exclude features with positive attributions. |
exclude_negative_attributions | bool | False | Exclude features with negative attributions. |
max_depth | int | 3 | Maximum depth of the treemap hierarchy. |
TreemapHierarchy
For custom treemap hierarchies:
from monad.interpretability import TreemapHierarchy
hierarchy = TreemapHierarchy(
levels=["category", "brand", "product_id"],
hierarchy_path=Path("./data/product_hierarchy.csv"),
feature_values_importance_path=Path("./interpretations/transactions/article_id/values_importance.json"),
entity_name_column="product_name", # Optional
)
generator = TreemapGenerator(hierarchy=hierarchy)
generator.plot_treemap(output_file_path=Path("./custom_treemap.html"))
| Field | Type | Default | Description |
|---|---|---|---|
levels | list[str] | required | Hierarchy level names for the treemap. |
hierarchy_path | Path | required | Path to a CSV file defining the hierarchy. Each row maps the last level to higher levels. |
feature_values_importance_path | Path | required | Path to the values importance JSON file. |
entity_name_column | str \| None | None | Column name for entity labels in the visualization. |
AttributionInterpreter and AttributionResults
monad.interpretability also exports get_interpreter(), AttributionInterpreter and AttributionResults, the building blocks interpret() and interpret_entity() run on. AttributionInterpreter.get_attributions(prediction_date, target, device) walks the predict set and returns an AttributionResults object with these fields:
| Field | Type | Description |
|---|---|---|
attributions | torch.Tensor | Attributions shaped (observations, features). |
feature_values | torch.Tensor | The inputs the attributions explain, with the same shape and row order. |
base_value | float | The model's average output over the reference rows. |
reference | str | Identifier of the reference the attributions are measured against, as recorded in run.json. |
get_interpreter() takes a run configuration object that monad.interpretability does not export. Use interpret() or interpret_entity() instead: they accept every run setting as a keyword argument, and interpret(..., save_shap_plots=True) renders the SHAP-style report from the same attributions.
Complete Workflow Example
from pathlib import Path
from datetime import datetime, timezone
from monad.ui.module import load_from_checkpoint
from monad.config import TestingParams, OutputType
from monad.interpretability import interpret, interpret_entity
from monad.interpretability import TreemapGenerator
# 1. Generate predictions with attributions enabled
module = load_from_checkpoint(Path("./churn_model"))
testing_params = TestingParams(
output_type=OutputType.SEMANTIC,
devices=[0],
prediction_date=datetime(2024, 6, 1, tzinfo=timezone.utc),
local_save_location=Path("./predictions.tsv"),
)
module.predict(testing_params)
# 2. Global feature importance
interpret(
predictions_path=Path("./predictions.tsv"),
output_path=Path("./interpretations"),
checkpoint_path=Path("./churn_model"),
device="cuda",
target_index=0, # binary classification
prediction_date=datetime(2024, 6, 1, tzinfo=timezone.utc),
)
# 3. Treemap visualization
generator = TreemapGenerator(
interpretability_files_path=Path("./interpretations"),
)
generator.plot_treemap(
output_file_path=Path("./treemap.html"),
)
# 4. Single entity deep-dive
interpret_entity(
output_path=Path("./interpretations/customer_123.json"),
checkpoint_path=Path("./churn_model"),
predictions_path=Path("./predictions.tsv"),
main_entity_id="customer_123",
device="cuda",
target_index=0,
prediction_date=datetime(2024, 6, 1, tzinfo=timezone.utc),
)
# 5. SHAP-style report for client hand-off
interpret(
predictions_path=Path("./predictions.tsv"),
output_path=Path("./interpretations"),
checkpoint_path=Path("./churn_model"),
device="cuda",
target_index=0,
method="gradient_shap",
save_shap_plots=True,
)