Skip to content

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 about steps_per_pass × batch_size rows, and neither the attributions nor the sample change.
  • batch_size: lower the number of rows per batch. With limit_batches set, this also shrinks the sample, because the run reads at most limit_batches × batch_size observations.
  • 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_value differ from runs made with those releases.
  • GradientSHAP samples up to n_baselines real observations at random, seeded by seed, 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.

from monad.interpretability import TreemapGenerator

Constructor

generator = TreemapGenerator(
    interpretability_files_path=Path("./interpretations"),
)
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,
)