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.