The architecture in context
The system we are building
A representation can learn from relationships between observations before receiving class labels. This experiment constructs two rotated views of each image and trains a shared CNN to keep matched views close relative to other images in the batch. A later ridge classifier measures how much useful class information the frozen representation retains.
Who does what in the stack
- PyTorch CNN
- Shared feature encoder for both views.
- PyTorch functional losses
- Normalized similarities and symmetric cross-entropy.
- Linear algebra solve
- Fits a fixed-feature ridge readout.
The archive changes the view-generation policy: nearby rotations stand in for temporal continuity, while independently sampled rotations provide an augmentation control. These are synthetic view pairs, not measured video trajectories or a demonstrated biological learning mechanism.
From module map to executable structure
Inside Contrastive image encoder
Rotating MNIST; grayscale input 28×28; embedding 128.
| Layer or branch | Output shape | Implementation detail |
|---|---|---|
| Two rotated views | 2 × B × 1 × 28 × 28 | Shared encoder weights; nearby versus independent angles are alternative data pipelines. |
| Conv 32 / BN / ReLU | B × 32 × 14 × 14 | 3×3 kernel, stride 2, padding 1. |
| Conv 64 / BN / ReLU | B × 64 × 7 × 7 | 3×3 kernel, stride 2, padding 1. |
| Conv 128 / BN / ReLU | B × 128 × 4 × 4 | 3×3 kernel, stride 2, padding 1. |
| Pool + embedding | B × 128 | Adaptive average 1×1 → flatten → Linear 128→128. Normalize embeddings inside the loss. |
The diagonal of the B×B cross-view similarity matrix contains positive pairs. Other columns are negatives for that row. Both directions contribute. Labels are not used for encoder training but are used in the downstream ridge probe. The two views must be different transforms of the same original digit, not independently sampled digit identities.
The equation and the update
Adam 1e-3; batch 256; temperature .2. Script default 20 epochs / temporal angle window 20°. E28’s saved comparison instead records 25 epochs / window 60° and a canonical probe. The distinction must survive reproduction.
Implementation card / no invented benchmarks
Capacity, budget and execution evidence
- Parameters / retained state
- 109,632 trainable scalars; BatchNorm running statistics are separate buffers.
- Duration and hardware evidence
- The saved JSON does not contain a verified end-to-end timing.
- Source coordinates
- E27 lines 23–25, 45–71 and 115–148
- Current reproduction context
- Current workstation, supplied by the author: Apple M4, 128 GB unified RAM, 40 GPU cores and 16 CPU cores. This is context for prospective reproduction, not attribution of every archived run. Python and framework versions are not fully locked for these historical sources; declarations, when available, are identified separately.
Counts above are calculated from the stated layer shapes unless identified as saved measurements. They exclude optimizer state and nontrainable buffers. No archived training was rerun for this revision.
What these design choices change
The channel expansion compensates for shrinking spatial resolution; global pooling discards spatial location before the probe. Temperature controls the sharpness of relative similarities. Increasing batch size changes both optimization and the number of negatives, so it is not purely a throughput change.
Reproduction and measurement protocol
Check pair indices before augmenting, compute the similarity matrix on a tiny batch, and verify that swapping both view orders together leaves the symmetric loss unchanged. Freeze the encoder and set eval mode before few-shot embedding extraction; otherwise BatchNorm updates contaminate the probe protocol.
For a new run, save the resolved Python/framework versions, backend, dtype, seed, input shapes, batch size and exact source revision. Start with one batch and one update. Log training steps separately from epochs or environment steps. Do not equate the configured maximum with a completed budget or convergence.
Measure initialization/compilation, data preparation, warmed forward pass, training updates and evaluation separately. Synchronize accelerator work around timed regions using the chosen framework’s supported mechanism. Report peak process memory and framework allocation separately; parameter bytes exclude activations, gradients, optimizer state and input buffers. On a shared machine, begin with a single CPU worker and a small batch rather than claiming all available resources.
A closer look at the implementation
The code that carries the idea
The snippet normalizes the embeddings, forms a cross-view similarity matrix and applies cross-entropy in both directions. The diagonal index supplies the positive pairing. Batch composition therefore defines the competing negatives; it is part of the loss, not just a throughput setting.
def info_nce(z1, z2, tau=0.2):
z1 = F.normalize(z1, dim=1); z2 = F.normalize(z2, dim=1)
logits = (z1 @ z2.T) / tau; labels = torch.arange(len(z1), device=z1.device)
return 0.5 * (F.cross_entropy(logits, labels) + F.cross_entropy(logits.T, labels))
def train_contrastive(enc, Xtr, mode, epochs, bs=256, seed=0, tstep=20.0):
opt = torch.optim.Adam(enc.parameters(), 1e-3); rng = torch.Generator(device="cpu"); rng.manual_seed(seed); N = len(Xtr)
for ep in range(epochs):
perm = torch.randperm(N, device=DEV)
for i in range(0, N, bs):
idx = perm[i:i + bs]; x = Xtr[idx]; B = len(x)
if mode == "temporal": # two nearby frames of a rotation sweep (window=tstep)
a0 = torch.rand(B, device=DEV) * 360.0; step = (torch.rand(B, device=DEV) * 2 - 1) * tstep
v1 = rotate(x, a0); v2 = rotate(x, a0 + step) # nearby viewpoints, same identity
else: # augment: two INDEPENDENT random rotations
v1 = rotate(x, torch.rand(B, device=DEV) * 360.0); v2 = rotate(x, torch.rand(B, device=DEV) * 360.0)
loss = info_nce(enc(v1), enc(v2)); opt.zero_grad(); loss.backward(); opt.step()
enc.eval(); return enc
Verbatim archive excerpt from closed_form_neat_invariance.py. Context-dependent historical code, not a standalone runnable program. Comments retain their original wording; the article distinguishes implemented behavior from stale or overbroad comments.
The boundary that matters
The learned invariance is only useful if the transformation preserves the intended label. Large rotations can make digit identity ambiguous, notably for some six/nine examples. Good pair discrimination does not by itself establish a good downstream classifier.