Skip to content

Recipe: LRP on a Trained MNIST Classifier

What you'll build: a Conv2DModule -> ReluModule -> FlattenModule -> LinearModule classifier (MnistConvNet) trained on MNIST digits. It is then explained with whole-model Layer-wise Relevance Propagation (LRP::explain), one held-out digit per class. The question is why the network predicts its class rather than the runner-up. Each digit is explained twice: with the epsilon rule everywhere, and with the EpsilonPlus composite.

CMake target: mnist_lrp_recipe (examples/recipes/mnist_lrp.cpp).

Run it from the repository root: ./build/mnist_lrp_recipe (Windows: build\Release\mnist_lrp_recipe.exe).

Note

Needs the MNIST IDX files in data/MNIST/raw/. Run python3 tools/fetch_mnist.py once first (the same requirement as examples/mnist_training_demo.cpp). Run the recipe from the repository root so the relative data path resolves. Training takes about 10 seconds on a CPU.

For the rules, composites and targets used here, see Layer-wise Relevance Propagation.

Code

MnistConvNet net(&backend);  // Conv2D(1,8,5,5) -> ReLU -> Flatten -> Linear(4608,10)
// ... net.train_step(image, label, optimizer, sink, step) over the training set ...

ExplainerContext ctx(net.modules());           // the trained layers, in forward order
const LRP epsilon;                             // epsilon rule on every layer
const LRP epsilon_plus = LRP::epsilon_plus();  // ZPlus on Conv2D, Zennit's epsilon on Linear

const Tensor x = batched(test.images[i]);      // (1, 1, 28, 28), so the output is (1, 10)

// "Why pred rather than runner_up?": target and contrast, seeded with the two logits.
const LRPTarget target{{pred}, {runner_up}};
const Attribution eps = epsilon.explain(ctx, x, target, &backend);
const Attribution eps_plus = epsilon_plus.explain(ctx, x, target, &backend);
// eps.values has the image's shape: one relevance value per pixel.
// eps.metadata.at("relevance_in_sum") vs. the logit margin shows how well relevance was conserved.

Full source: examples/recipes/mnist_lrp.cpp.

Expected output

MNIST LRP recipe -- MnistConvNet: Conv2D(1,8,5,5)->ReLU->Flatten->Linear(4608,10), Adam(lr=0.001)
Training on 10000 real images, evaluating on 1000 held-out images, 3 epochs

epoch 1 | mean train loss 0.2653 | test accuracy 95.6%
epoch 2 | mean train loss 0.1039 | test accuracy 95.4%
epoch 3 | mean train loss 0.0590 | test accuracy 96.0%

LRP on the first held-out example of each digit: why the prediction rather than the runner-up
(target = prediction, contrast = runner-up class, seeded with the two logits)

label  pred  runner-up   margin   sum(R) eps  sum(R) eps-plus
7      7     2           17.466    17.363594        17.264297
2      2     6            2.737     3.723733         3.578975
1      1     8            5.609     6.890915         7.359929
0      0     6            9.759     9.837962         9.564141
4      4     9            8.495     9.635921         9.480172
9      9     4            6.131     5.523613         5.190765
5      6     5            2.697     2.386827         2.447210   <- misclassified
6      6     0            1.210     1.283611         1.300158
3      8     5            0.648     0.826104         0.532464   <- misclassified
8      8     2           13.095    13.480010        12.657933
...

The recipe then prints the relevance maps for the first held-out 0 as ASCII, epsilon and EpsilonPlus side by side (+ is evidence for 0, - is evidence for the runner-up).

What's happening

Why a contrastive target instead of the raw logit. LRP redistributes whatever score you seed it with. A trained logit can be positive only because of its bias. The epsilon rule redistributes the pre-bias sum, so seeding the raw winning logit (or a one-hot seed of 1) can flip the sign of the whole heatmap.

LRPTarget{{pred}, {runner_up}} asks a more useful question: "why this class rather than the next best one?" LRP is linear in the seed, so the result is exactly explanation(pred) − explanation(runner-up). The explained score is the logit margin, which is positive whenever the prediction wins. See Targets, contrasts and seeds.

Rule choice. LRP() applies the epsilon rule on every layer. LRP::epsilon_plus() is the EpsilonPlus composite from Zennit (a PyTorch LRP library). It uses ZPlus on the convolution, which keeps only positive contributions and gives a less noisy map. On the classifier it uses epsilon with the bias in the denominator. See Composites for the rule definitions.

Conservation. sum(R) over the 784 pixels tracks the margin, but not exactly. The gap is largest when the margin is small. Under epsilon, the classifier, ReLU and Flatten conserve relevance exactly. The whole gap comes from Conv2DModule units in the blank background around the digit. Their input patch is all zeros, so they're active only because of their bias. They still receive relevance from the classifier, but with zero input there's nothing to pass it on to, so it's dropped. That dropped relevance can be positive or negative, which is why sum(R) lands on either side of the margin. On the first four digits (7, 2, 1 and 0), the gap equals the relevance on those units to four decimals.

EpsilonPlus also keeps the classifier's bias in the denominator, so the bias absorbs its share of the margin as well.

The misclassified digits. The 5 that the network calls a 6 has the true class as its runner-up. Its contrastive map shows which strokes tipped it from 5 to 6 (margin 2.7). The 3 called an 8 is a near tie with 5 (margin 0.65), and its true class isn't even second. To ask "why 8 rather than 3", pass the true label as the contrast instead of the runner-up.

See also: Grad-CAM walkthrough for the same layer types (Conv2D → ReLU → Flatten → Linear) explained with gradients on an untrained toy network, and the MNIST gallery for every explainer on this model rendered as heatmaps.