Skip to content

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/knn
overrides dict_metric_from_list()

Two things that sound hard here are not:

  • A dataset returning two views. prep_batch maps 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 no tgt_name at 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:

  1. A task that is too easy. At style_scale = 1.5, valid/knn reaches 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 uses style_scale = 4.0.
  2. 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.
  3. 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.