TorchXAI wraps Captum attribution methods and adds first-class multi-target support: explain multiple output targets in a single forward pass.

Available Explainers

Base Classes

Gradient-Based Methods

Perturbation-Based Methods

  • Feature Ablation — systematically zeros out features or groups
  • Occlusion — sliding-window patch replacement
  • LIME — locally-linear surrogate model
  • Kernel SHAP — Shapley values via LIME kernel weighting

Baseline Methods

  • Random — random attributions for sanity-checking and comparison

Input Patterns

Each explainer belongs to one of five input patterns. Choose the pattern for your explainer and your use case:

Pattern Required arguments Explainers
A inputs, target Saliency, InputXGradient, GuidedBackprop, Random
B inputs, baselines, target IntegratedGradients, DeepLift, InputXBaselineGradient
C inputs, baselines (distribution), target GradientShap, DeepLiftShap
D inputs, feature_mask (optional), target FeatureAblation, LIME, KernelShap
E inputs, sliding_window_shapes, target Occlusion

Full comparison table

Explainer Type baselines Baseline distribution feature_mask sliding_window_shapes
SaliencyExplainer Gradient
InputXGradientExplainer Gradient
GuidedBackpropExplainer Gradient
RandomExplainer Baseline
IntegratedGradientsExplainer Gradient
DeepLiftExplainer Gradient
InputXBaselineGradientExplainer Gradient
DeepLiftShapExplainer Gradient
GradientShapExplainer Gradient
FeatureAblationExplainer Perturbation optional
LimeExplainer Perturbation optional
KernelShapExplainer Perturbation optional
OcclusionExplainer Perturbation

Quick Start

import torch
import torch.nn as nn
from torchxai.explainers import SaliencyExplainer, IntegratedGradientsExplainer
from torchxai.data_types import SingleTargetAcrossBatch

model = nn.Sequential(nn.Linear(10, 5), nn.ReLU(), nn.Linear(5, 3))
model.eval()

inputs   = torch.randn(1, 10)
baseline = torch.zeros(1, 10)
target   = SingleTargetAcrossBatch(index=0)   # explain class 0

# Pattern A — no baseline needed
explainer = SaliencyExplainer(model)
attrs = explainer.explain(inputs=inputs, target=target)
print(attrs.shape)   # (1, 10)

# Pattern B — baseline required
explainer_ig = IntegratedGradientsExplainer(model)
attrs_ig = explainer_ig.explain(inputs=inputs, baselines=baseline, target=target)
print(attrs_ig.shape)   # (1, 10)

Multi-Target Mode

Pass multi_target=True at construction, then supply a list of targets. The explainer returns a list[Tensor], one per target, in a single forward-backward pass.

from torchxai.explainers import SaliencyExplainer
from torchxai.data_types import SingleTargetAcrossBatch

targets = [SingleTargetAcrossBatch(index=i) for i in range(3)]   # classes 0, 1, 2

explainer = SaliencyExplainer(model, multi_target=True)
attrs_list = explainer.explain(inputs=inputs, target=targets)

for cls_idx, attr in enumerate(attrs_list):
    print(f"class {cls_idx}: {attr.shape}")   # each (1, 10)

This is equivalent to calling explain() once per target but can be significantly faster because shared computation (forward pass, intermediate activations) is reused.


End-to-End Examples

Worked examples covering all five input patterns:


Sanity-Checking with Random Attributions

RandomExplainer provides random-noise attributions. Use it to verify that your real explainer is producing signal above chance:

from torchxai.explainers import RandomExplainer, SaliencyExplainer
from torchxai.data_types import SingleTargetAcrossBatch

target = SingleTargetAcrossBatch(index=0)

random_attrs   = RandomExplainer(model, random_seed=42).explain(inputs=inputs, target=target)
saliency_attrs = SaliencyExplainer(model).explain(inputs=inputs, target=target)

print("Random  :", random_attrs.abs().mean().item())
print("Saliency:", saliency_attrs.abs().mean().item())

If saliency attribution magnitudes are comparable to random, the model is likely not using that feature meaningfully — or the explainer is misconfigured.