Add an online probe

Add an online probe#

An OnlineProbe trains a classifier on detached features. The encoder continues to optimize its SSL objective. The forward function should return embedding and, when views duplicate samples, a correspondingly repeated label tensor.

This complete construction can be adapted to the training script from the quickstart:

import torch
from torchmetrics.classification import MulticlassAccuracy
import stable_pretraining as spt

module = spt.Module(
    forward=spt.forward.simclr,
    backbone=torch.nn.Sequential(torch.nn.Flatten(), torch.nn.Linear(3 * 16 * 16, 16)),
    projector=torch.nn.Linear(16, 8),
    simclr_loss=spt.losses.NTXEntLoss(temperature=0.5),
    optim={"optimizer": {"type": "AdamW", "lr": 1e-3}},
)
probe = spt.OnlineProbe(
    module,
    name="probe",
    input="embedding",
    target="label",
    probe=torch.nn.Linear(16, 2),
    loss=torch.nn.CrossEntropyLoss(),
    metrics={"accuracy": MulticlassAccuracy(2)},
)

Pass callbacks=[probe] to lightning.Trainer and run through spt.Manager. The quickstart already does this; its eval/probe_accuracy metric comes from held-out validation images. Feature width must match the probe’s input width, and prediction and label batch lengths must match. An online probe is a training diagnostic; report a separate, controlled frozen-encoder evaluation when your research protocol requires it.

The regression suite compares backbone updates with and without evaluation callbacks. A probe should not change the number of backbone optimizer steps.