Evidential Deep Learning for Classification Uncertainty: A Complete Implementation Guide

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 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. 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.


# 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 within each loss module:

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 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.

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 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.


# 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), 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 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:


# run_edl.py

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

if __name__ == '__main__':
    main()

Execute with:

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:

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:

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.
  • 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 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. 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 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 and configured through Configs.kl_div_coef.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →