Masksembles on MNIST

Replace a full ensemble with a fixed set of binary masks inserted after each hidden layer of a shared backbone. During training one mask is drawn per sample (dropout-style); during inference the model runs once per mask.

from __future__ import annotations

import numpy as np
import torch
from torch import nn

from probly.predictor import predict
from probly.quantification import quantify
from probly.representer import representer
from probly.transformation.masksembles import masksembles
from probly_benchmark.data import load_mnist

from examples.utils.model import MLPClassifier
from examples.utils.plotting import plot_mnist_uncertainty

Setup

train_loader, test_loader = load_mnist(batch_size=64)

X_test_batches, y_test_batches = zip(*test_loader)
X_test = torch.cat([x.view(-1, 28 * 28) for x in X_test_batches])
y_test = torch.cat(list(y_test_batches))
images_test = (X_test.view(-1, 28, 28) * 255).byte()

Model

num_masks binary masks are generated and inserted after each linear layer. A larger scale reduces overlap (and thus correlation) between masks at the cost of capacity per masked sub-network.

base_model = MLPClassifier(in_features=28 * 28, hidden_features=256, out_features=10)
num_masks = 4
masksembles_model = masksembles(
    base_model,
    num_masks=num_masks,  # The higher, the more similar to MC Dropout
    scale=2.0,  # The higher, the more similar to Ensemble
    predictor_type="logit_classifier",
)

print(masksembles_model)
MLPClassifier(
  (net): Sequential(
    (0_0): Linear(in_features=784, out_features=256, bias=True)
    (0_1): MasksemblesLinear()
    (1): ReLU()
    (2_0): Linear(in_features=256, out_features=256, bias=True)
    (2_1): MasksemblesLinear()
    (3): ReLU()
    (4): Linear(in_features=256, out_features=10, bias=True)
  )
)

Training

Fine-tune with standard cross-entropy. In training mode each sample in the batch is masked independently with a uniformly random mask (dropout-style); no manual batch tiling is needed.

opt = torch.optim.Adam(masksembles_model.parameters(), lr=1e-4)

epochs = 5

masksembles_model.train()
for epoch in range(epochs):
    train_loss, train_correct = 0, 0
    for x_batch, y_batch in train_loader:
        X_flat = x_batch.view(-1, 28 * 28)

        opt.zero_grad()
        out = masksembles_model(X_flat)
        loss = nn.functional.cross_entropy(out, y_batch)
        loss.backward()
        opt.step()

        train_loss += loss.item() * X_flat.size(0)
        train_correct += (out.argmax(1) == y_batch).sum().item()

    # Validation: ``predict`` tiles each batch by ``num_masks`` and returns one
    # prediction per mask; averaging the per-mask probabilities gives the
    # ensemble prediction. (A direct eval-mode forward would instead expect a
    # batch that is already tiled by ``num_masks``.)
    val_loss, val_correct = 0, 0
    with torch.no_grad():
        for x_batch, y_batch in test_loader:
            X_flat = x_batch.view(-1, 28 * 28)
            mean_probs = predict(masksembles_model, X_flat).tensor.softmax(-1).mean(0)
            loss = nn.functional.nll_loss(mean_probs.log(), y_batch)
            val_loss += loss.item() * X_flat.size(0)
            val_correct += (mean_probs.argmax(1) == y_batch).sum().item()

    print(
        f"Epoch {epoch+1}/{epochs} "
        f"- Train loss: {train_loss/len(train_loader.dataset):.4f}, "
        f"Train acc: {train_correct/len(train_loader.dataset):.4f}, "
        f"Val loss: {val_loss/len(test_loader.dataset):.4f}, "
        f"Val acc: {val_correct/len(test_loader.dataset):.4f}"
    )
Epoch 1/5 - Train loss: 0.9771, Train acc: 0.7553, Val loss: 0.4125, Val acc: 0.8880
Epoch 2/5 - Train loss: 0.3778, Train acc: 0.8936, Val loss: 0.3228, Val acc: 0.9092
Epoch 3/5 - Train loss: 0.3193, Train acc: 0.9093, Val loss: 0.2871, Val acc: 0.9169
Epoch 4/5 - Train loss: 0.2888, Train acc: 0.9181, Val loss: 0.2626, Val acc: 0.9244
Epoch 5/5 - Train loss: 0.2656, Train acc: 0.9238, Val loss: 0.2475, Val acc: 0.9283

Uncertainty Quantification

masksembles_model.eval()
rep = representer(masksembles_model)

with torch.no_grad():
    representation = rep.represent(X_test)

uq = quantify(representation)
_total = uq.total
uncertainty = (
    _total.detach().numpy() if isinstance(_total, torch.Tensor) else np.asarray(_total)
)
uncertainty = uncertainty / np.log(2)
if uncertainty.ndim > 1:
    uncertainty = uncertainty.sum(axis=-1)

Predictions

predict tiles the input by num_masks internally and returns a TorchSample of shape [num_masks, N, num_classes] — one slice per mask. Softmax converts logits to probabilities; averaging over masks gives the mean predictive distribution used for the final class prediction.

with torch.no_grad():
    sample = predict(masksembles_model, X_test)              # TorchSample [num_masks, N, 10]

member_probs = sample.tensor.softmax(-1).numpy()             # [num_masks, N, 10]
mean_probs = member_probs.mean(axis=0)                       # [N, 10]

accuracy = (mean_probs.argmax(-1) == y_test.numpy()).mean() * 100
print(f"Test accuracy: {accuracy:.1f}%")
Test accuracy: 92.8%

Visualization

Plot the five most uncertain test digits with per-member agreement.

plot = plot_mnist_uncertainty(
    images_test,
    y_test,
    uncertainty,
    mean_probs,
    title="Top-5 Most Uncertain Test Predictions (Masksembles)",
)
plot.show()
Top-5 Most Uncertain Test Predictions (Masksembles), True: 4 | Pred: 6 U = 3.04 bits, True: 0 | Pred: 2 U = 2.86 bits, True: 5 | Pred: 5 U = 2.84 bits, True: 1 | Pred: 1 U = 2.77 bits, True: 8 | Pred: 5 U = 2.71 bits

Total running time of the script: (1 minutes 2.627 seconds)

Gallery generated by Sphinx-Gallery