# Evidential Deep Learning for Classification Uncertainty: A Complete Implementation Guide

> Learn to implement evidential deep learning for classification uncertainty. Transform neural network outputs into evidence parameters defining a Dirichlet distribution for explicit uncertainty quantification with labmlai.

- Repository: [labml.ai/annotated_deep_learning_paper_implementations](https://github.com/labmlai/annotated_deep_learning_paper_implementations)
- Tags: how-to-guide
- Published: 2026-03-04

---

**The labmlai/annotated_deep_learning_paper_implementations repository implements evidential deep learning (EDL) by transforming neural network outputs into non-negative evidence parameters that define a Dirichlet distribution, enabling explicit uncertainty quantification through specialized Bayes risk losses and KL regularization.**

This guide walks through the complete implementation of **evidential deep learning for classification uncertainty** in the [labmlai/annotated_deep_learning_paper_implementations](https://github.com/labmlai/annotated_deep_learning_paper_implementations) repository. The code demonstrates how to replace traditional softmax probability estimates with Dirichlet-based uncertainty modeling on the MNIST dataset.

## Model Architecture and Evidence Extraction

The implementation uses a standard LeNet-style CNN defined in [`labml_nn/uncertainty/evidence/experiment.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/uncertainty/evidence/experiment.py). The `Model` class outputs raw logits through a final linear layer (`self.fc2`) producing a tensor of shape `[batch_size, 10]` for the ten MNIST classes.

To ensure non-negative evidence values required for Dirichlet parameterization, the raw outputs pass through an activation function selected via the `outputs_to_evidence` configuration. The repository supports both **ReLU** and **Softplus** transformations.

```python

# From labml_nn/uncertainty/evidence/experiment.py

outputs = self.model(data)                     # Raw logits [batch, 10]

evidence = self.outputs_to_evidence(outputs)   # Non-negative evidence e_k ≥ 0

```

The evidence transformation is configured in the `Configs` class, allowing flexible switching between activation functions without modifying the model architecture.

## Dirichlet Parameterization

The core mathematical operation converts evidence into **Dirichlet concentration parameters** (alphas). For each class $k$, the implementation calculates:

$$\alpha_k = e_k + 1$$

This transformation appears in [`labml_nn/uncertainty/evidence/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/uncertainty/evidence/__init__.py) within each loss module:

```python
alpha = evidence + 1.

```

The concentration parameters $\alpha$ define a Dirichlet distribution $\text{Dir}(p|\alpha)$ over the class probability simplex, where the strength of evidence directly influences the sharpness of the distribution. Higher evidence values produce more confident (peaked) distributions, while zero evidence yields a uniform distribution representing maximum uncertainty.

## Evidential Loss Functions

The repository provides four complementary loss components in [`labml_nn/uncertainty/evidence/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/uncertainty/evidence/__init__.py) to train the model:

**MaximumLikelihoodLoss** implements the negative log-marginal likelihood of the Dirichlet prior (Type-II maximum likelihood). This loss maximizes the probability of observed data under the Dirichlet distribution.

**CrossEntropyBayesRisk** computes the Bayes risk using cross-entropy as the cost function. The implementation utilizes digamma functions to calculate the expected log-probabilities under the Dirichlet distribution efficiently.

**SquaredErrorBayesRisk** provides an alternative Bayes risk formulation using squared-error cost. This loss explicitly decomposes into error and variance terms, offering interpretable gradients during training.

**KLDivergenceLoss** serves as a regularizer that pulls the Dirichlet distribution toward a uniform prior for misclassified samples. This prevents overconfident predictions on out-of-distribution data.

```python
from labml_nn.uncertainty.evidence import (
    MaximumLikelihoodLoss,
    CrossEntropyBayesRisk,
    SquaredErrorBayesRisk,
    KLDivergenceLoss
)

# Example usage

evidence = torch.abs(torch.randn(32, 10))  # Non-negative evidence

target = torch.eye(10)[torch.randint(0, 10, (32,))]

loss_fn = CrossEntropyBayesRisk()
loss = loss_fn(evidence, target)

```

## Training Loop with KL Annealing

The training procedure in [`experiment.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/experiment.py) combines a primary Bayes risk loss with the KL divergence regularizer. A critical implementation detail is the **annealing coefficient** that gradually increases the strength of KL regularization during training.

```python

# From experiment.py step method (lines 33-50)

outputs = self.model(data)
evidence = self.outputs_to_evidence(outputs)
loss = self.loss_func(evidence, target)
kl_div_loss = self.kl_div_loss(evidence, target)

# Annealing schedule prevents early over-regularization

annealing_coef = min(1., self.kl_div_coef(tracker.get_global_step()))
total_loss = loss + annealing_coef * kl_div_loss

```

The annealing coefficient increases from 0 to 1 according to a configurable schedule (defined in [`labml_nn/helpers/schedule.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/helpers/schedule.py)), ensuring the model first learns to fit the data before strong regularization constraints take effect.

## Uncertainty Quantification Metrics

The `TrackStatistics` class in [`labml_nn/uncertainty/evidence/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/uncertainty/evidence/__init__.py) extracts interpretable uncertainty metrics from the Dirichlet distribution:

- **Uncertainty mass**: $u = K / S$, where $S = \sum_k \alpha_k$ represents the total evidence strength
- **Expected class probability**: $\hat{p}_k = \alpha_k / S$
- **Variance**: Captures epistemic uncertainty about the probability estimates themselves

These metrics are logged via LabML's tracker system for both correct and incorrect predictions, enabling analysis of uncertainty calibration across different confidence levels.

## Practical Implementation Examples

### Running the MNIST Experiment

To train the evidential model on MNIST:

```python

# run_edl.py

from labml import experiment
from labml_nn.uncertainty.evidence.experiment import main

if __name__ == '__main__':
    main()

```

Execute with:

```bash
python run_edl.py

```

Key configuration options in the `Configs` class include:
- `loss_func`: Select `'max_likelihood_loss'`, `'cross_entropy_bayes_risk'`, or `'squared_error_bayes_risk'`
- `outputs_to_evidence`: Choose `'relu'` or `'softplus'`
- `kl_div_coef_schedule`: Control the KL annealing curve

### Custom Model Integration

Replace the default LeNet backbone while preserving evidential functionality:

```python
from labml_nn.uncertainty.evidence.experiment import Configs
from labml import experiment
import torch.nn as nn
from labml.configs import option

class CustomCNN(nn.Module):
    def __init__(self, dropout: float = 0.5):
        super().__init__()
        self.conv = nn.Sequential(
            nn.Conv2d(1, 32, 3), nn.ReLU(), nn.MaxPool2d(2),
            nn.Conv2d(32, 64, 3), nn.ReLU(), nn.MaxPool2d(2)
        )
        self.fc = nn.Linear(64 * 5 * 5, 10)
        self.dropout = nn.Dropout(dropout)
    
    def forward(self, x):
        x = self.conv(x).view(x.size(0), -1)
        return self.fc(self.dropout(x))

@option(Configs.model)
def custom_model(c: Configs):
    return CustomCNN(c.dropout).to(c.device)

if __name__ == '__main__':
    experiment.create(name='custom_edl')
    Configs().run()

```

### Direct Loss Module Usage

Integrate evidential losses into existing training pipelines:

```python
import torch
from labml_nn.uncertainty.evidence import (
    SquaredErrorBayesRisk, 
    KLDivergenceLoss,
    TrackStatistics
)

evidence = torch.abs(torch.randn(16, 10))
targets = torch.randint(0, 10, (16,))

# Compute losses

primary_loss = SquaredErrorBayesRisk()(evidence, targets)
kl_loss = KLDivergenceLoss()(evidence, targets)

# Monitor statistics

stats = TrackStatistics()
stats(evidence, targets)  # Logs uncertainty mass and accuracy

```

## Summary

- **Evidence transformation**: Raw logits become non-negative evidence via ReLU or Softplus activations in [`labml_nn/uncertainty/evidence/experiment.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/uncertainty/evidence/experiment.py).
- **Dirichlet construction**: Evidence converts to concentration parameters $\alpha = e + 1$, defining a distribution over class probabilities.
- **Multi-component loss**: The implementation combines Bayes risk losses (cross-entropy or squared error) with KL divergence regularization toward a uniform prior.
- **Annealing strategy**: KL regularization strength increases gradually during training to prevent early over-constraint.
- **Explicit uncertainty**: The `TrackStatistics` class provides uncertainty mass and expected probabilities derived from Dirichlet parameters.

## Frequently Asked Questions

### What is evidential deep learning?

Evidential deep learning is a framework that treats neural network outputs as evidence for a Dirichlet distribution over class probabilities rather than point estimates. This approach models second-order uncertainty (uncertainty about the probabilities themselves), allowing the model to express "I don't know" when evidence is scarce. The labmlai implementation achieves this by constraining outputs to non-negative values and interpreting them as pseudo-counts in a Dirichlet distribution.

### How does EDL quantify classification uncertainty differently than softmax?

Traditional softmax outputs probabilities that sum to 1.0, forcing the model to commit to some distribution even when uncertain. In contrast, EDL in [`labml_nn/uncertainty/evidence/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/uncertainty/evidence/__init__.py) uses the total evidence $S = \sum_k \alpha_k$ to compute uncertainty mass $u = K/S$. When evidence is low, $u$ approaches 1.0 (maximum uncertainty), while high evidence drives $u$ toward 0. This allows the model to distinguish between high-confidence predictions (sharp Dirichlet) and ambiguous inputs (flat Dirichlet).

### Which loss function should I use for evidential classification?

The repository provides three primary options in [`labml_nn/uncertainty/evidence/__init__.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/uncertainty/evidence/__init__.py). **MaximumLikelihoodLoss** works well for standard supervised learning. **CrossEntropyBayesRisk** often provides better calibration for classification tasks. **SquaredErrorBayesRisk** offers explicit variance decomposition useful when you need to separate aleatoric and epistemic uncertainty. All three benefit from pairing with **KLDivergenceLoss** for regularization.

### Why is KL annealing necessary in evidential deep learning?

Without annealing, the KL divergence loss in [`experiment.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/experiment.py) would immediately force the Dirichlet toward a uniform prior, preventing the model from learning meaningful patterns. The annealing coefficient (`annealing_coef`) starts near zero and increases to 1.0 over training steps, allowing the model to first accumulate evidence for correct classifications before regularizing uncertain predictions. This schedule is implemented via [`labml_nn/helpers/schedule.py`](https://github.com/labmlai/annotated_deep_learning_paper_implementations/blob/main/labml_nn/helpers/schedule.py) and configured through `Configs.kl_div_coef`.