Note
Go to the end to download the full example code.
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()

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