Contrastive / self-supervised (SimCLR)¶
Two augmented views per sample, NT-Xent loss, kNN probe on frozen features. CPU, no downloads, ~15 s.
gatle-ignite train --config=examples/contrastive/configs/contrastive_v0.py
python examples/contrastive/scripts/probe_check.py # the proof
Nothing beyond the walk: dataset · model · optimizer · checkpoints · logging · running
| step | what this example changes |
|---|---|
| prep_batch | overrides build_dataloaders() |
| loss | dict_of_loss_params: ntxent (ntxent) |
| metrics | score_name = valid/knnoverrides dict_metric_from_list() |
Two things that sound hard here are not:
- A dataset returning two views.
prep_batchmaps them to two model inputs. No special support. - A loss with no labels. NT-Xent's targets are positions in the batch, so the sub-loss simply
reads two
src_*names and notgt_nameat all. Nothing requires a loss to have a target.
examples/contrastive/losses/loss_functions/ntxent.py
import torch
import torch.nn as nn
import torch.nn.functional as F
from gatle_ignite import get_value
class Loss(nn.Module):
def __init__(self, src_a="z1", src_b="z2", temperature=0.2):
super().__init__()
# Two fields, not a tuple: get_value's tuple form is a PATH into nested dicts.
self.src_a, self.src_b = src_a, src_b
self.temperature = temperature
def forward(self, y_pred, target=None, iteration=None):
z1 = get_value(y_pred, self.src_a) # (B, D), already L2-normalised by the model
z2 = get_value(y_pred, self.src_b) # (B, D)
b = z1.shape[0]
if b < 2:
raise ValueError(
f"NT-Xent needs at least 2 samples per batch, got {b}: with one sample "
"there are no negatives and the loss is identically zero."
)
z = torch.cat([z1, z2], dim=0) # (2B, D)
sim = (z @ z.T) / self.temperature # (2B, 2B)
# Mask the diagonal with -inf: a sample is not its own negative. Out-of-place, since
# in-place on an autograd-tracked output risks a version-counter error.
eye = torch.eye(2 * b, dtype=torch.bool, device=z.device)
sim = sim.masked_fill(eye, float("-inf"))
# THE TARGETS: i's positive is at i+B and (i+B)'s at i, indices into this batch.
targets = torch.cat(
[torch.arange(b, 2 * b, device=z.device), torch.arange(0, b, device=z.device)]
)
return F.cross_entropy(sim, targets)
Why the loss is not the evidence¶
NT-Xent falls for degenerate representations too. The evidence is a kNN probe on frozen features,
against labels the loss never sees, built imperatively via dict_metric_from_list, because it has
to close over the live model.
chance : 0.1000
raw inputs (no encoder) : 0.4004
untrained encoder (control) : 0.3643
trained encoder : 0.9756
verdict: LEARNED
The untrained encoder scoring below raw inputs (0.364 vs 0.400) is what a working control should
show: a random projection destroys information. probe_check.py exits non-zero on NO EVIDENCE.
Three ways an SSL result can fool you¶
Worth reading if you are evaluating SSL, because each one looks like success:
- A task that is too easy. At
style_scale = 1.5,valid/knnreaches 1.0 at epoch 1, but raw inputs and the untrained encoder score 1.000 too: a falling loss and 100% accuracy, proving nothing. This example usesstyle_scale = 4.0. - A control that loads the trained weights. Inference mode evaluates what is on disk, not what
is in memory (see
evaluate()), so an "untrained" control run through it loads the trained checkpoint and scores the same as the trained model. That reads as a plausible "SSL didn't work" result, when the numbers being identical is the clue. - Random draws in the collate. They make the eval set non-deterministic, so one checkpoint scores differently across runs. Every random draw here happens in the dataset.