Quickstart
1. Get model outputs and labels
You need two things from a held-out, labelled calibration set (see Concepts for why it must be held out):
predictions: the model output, shape(B, C, *spatial)with the class axis at dimension 1 (e.g.(B, C, H, W)in 2D,(B, C, D, H, W)in 3D);targets: integer labels, shape(B, *spatial).
import torch
import torch.nn.functional as F
from fiducio import TemperatureScaling, negative_log_likelihood, expected_calibration_error
# Stand-in for your model's outputs on the calibration set.
logits = torch.randn(8, 4, 64, 64)
labels = torch.randint(0, 4, (8, 64, 64))
2. Fit a calibrator
3. Apply it to new predictions
new_logits = torch.randn(2, 4, 64, 64)
calibrated = calibrator.transform(new_logits) # calibrated probabilities (B, C, H, W)
transform always returns probabilities that sum to 1 along the class axis.
predict_proba is an alias, and fit_transform(logits, labels) does both steps.
If you need calibrated logits instead, use decision_function(new_logits)
(softmax of its output equals transform).
4. Measure the effect
# Separate test data, never used for fitting or model selection:
test_logits = torch.randn(8, 4, 64, 64)
test_labels = torch.randint(0, 4, (8, 64, 64))
raw = F.softmax(test_logits, dim=1)
calibrated_test = calibrator.transform(test_logits)
print("NLL", negative_log_likelihood(raw, test_labels), "->",
negative_log_likelihood(calibrated_test, test_labels))
print("ECE", expected_calibration_error(raw, test_labels), "->",
expected_calibration_error(calibrated_test, test_labels))
5. Save and reload
from fiducio import load_calibrator
calibrator.save("calibrator.pt")
calibrator = load_calibrator("calibrator.pt") # loads on CPU by default
Logits or probabilities?
Set input_type to match what you pass in:
TemperatureScaling(input_type="logits") # raw scores (default)
TemperatureScaling(input_type="probs") # probabilities that sum to 1
Accepted tensor shapes
| Use case | predictions |
targets |
|---|---|---|
| 2D segmentation | (B, C, H, W) |
(B, H, W) |
| 3D segmentation | (B, C, D, H, W) |
(B, D, H, W) |
| general n-D | (B, C, *spatial) |
(B, *spatial) |
| per-pixel table | (N, C) |
(N,) |
The class axis is always dimension 1. Binary segmentation is provided as two
channels (C = 2).