Skip to content
AITroveRead. Build. Understand.
Make this comfortable

Occlusion attribution and explanation sanity checks

Last updated: 6 Oct 20265 min read
tutorial
AdvancedBy AITrove Editorial

Replacing one input region and measuring a score change can locate model sensitivity, but the map depends on the replacement and does not prove causal human reasoning.

Define the score and replacement

Fix a target class logit, then replace one image patch at a time with a stated baseline such as a neutral local intensity. The difference between the original and altered logit is a patch sensitivity under that intervention. A zero patch on a white receipt creates an artificial black square; a local-median patch has a different meaning. Use more than one plausible baseline and note where results disagree. Spatial geometry determines which input region maps to a feature response.

Keep preprocessing identical

The image supplied to an explanation probe must pass through the same resize, normalization and channel order as ordinary inference. Apply a replacement in a declared space: raw pixels before normalization or normalized tensors after it. Mixing those spaces can create an out-of-range patch and a large score difference unrelated to the model’s ordinary evidence. Use eval mode and a frozen model revision. Log the exact target class; a map for the currently predicted class can change target when the patch is removed.

Inspect rather than certify

A bright patch map may overlap the printed receipt defect, yet that alone does not establish why the classifier made its decision. Test an untrained or randomized-weight model with the same probe, compare maps after controlled input changes, and measure whether removing top-ranked patches changes the target score more than removing equally sized random patches. These tests may expose maps that mainly reflect edges or preprocessing. A localization label provides independent evidence for overlap checks.

Budget the probe

For an image partitioned into R patches, a simple occlusion map uses one original forward pass plus R altered forward passes. That can be expensive for large images or an online endpoint. Batch altered inputs if memory permits, and run the full audit offline if it would violate serving latency. A coarse patch grid is cheaper but can miss small evidence; overlapping small patches cost more and spread one pixel across several interventions. Report patch size and stride with every map.

Use failure slices

Check low-contrast, edge-truncated and correctly classified receipts separately from high-confidence mistakes. A map that follows a scanner border instead of text may reveal shortcut learning. If removing a region flips the class, record the changed score and whether the altered image remains plausible. Do not turn an attractive heatmap into a user-facing claim of proof. Confidence and attribution answer different questions and both need held-out checks.

Implementation

python
import torch
from torch import nn

torch.manual_seed(89)
receipt_patch = torch.tensor([[[[0.9, 0.8, 0.7, 0.9],
                                [0.9, 0.2, 0.3, 0.8],
                                [0.8, 0.1, 0.2, 0.7],
                                [0.9, 0.8, 0.9, 0.8]]]])
quality_model = nn.Sequential(nn.Flatten(), nn.Linear(16, 2)).eval()
randomized_model = nn.Sequential(nn.Flatten(), nn.Linear(16, 2)).eval()

def patch_sensitivity(model: nn.Module, image: torch.Tensor,
                      target_class: int, replacement: float) -> torch.Tensor:
    with torch.inference_mode():
        original_score = model(image)[0, target_class]
        score_changes = torch.zeros(2, 2)
        for patch_row in range(2):
            for patch_column in range(2):
                altered = image.clone()
                top, left = patch_row * 2, patch_column * 2
                altered[:, :, top:top + 2, left:left + 2] = replacement
                score_changes[patch_row, patch_column] = (
                    original_score - model(altered)[0, target_class])
    return score_changes

attribution = patch_sensitivity(quality_model, receipt_patch, 1, 0.8)
randomized_attribution = patch_sensitivity(randomized_model, receipt_patch, 1, 0.8)
parameter_sensitivity = (attribution - randomized_attribution).abs().mean()
assert attribution.shape == (2, 2)
assert torch.isfinite(parameter_sensitivity)

Performance and operating cost

A naive patch probe requires O(R × F) model work for R regions and one forward cost F, plus storage for a small map; batching probes trades memory for throughput. Comparison with a randomized model roughly doubles that work. The interpretation also has a data cost: patch replacement can create unrealistic images. Measure probe stability under several baselines and keep ordinary inference independent of the optional explanation pass.

Common Mistakes

  • Do not call a sensitivity map causal proof of a human-readable reason.
  • Do not replace pixels in a different normalization space than the model expects.
  • Do not judge a map only by visual appeal without model and input counterfactuals.

Read next

Continue the workflow: FGSM input scale and perturbation threat models.

Continue the workflow: Graph message passing, edge direction and self state.

ai-data
deep-learning
Storage details