CLIP learns a shared embedding space for images and text by pulling matching image-caption pairs together and pushing every non-matching pair apart. The idea is simple, but the systems problem is unusual. In most training, batch size is a throughput knob you tune for hardware efficiency. In contrastive training, the batch is also the set of negatives every example is compared against, so batch size directly changes what the model learns. The original CLIP used a global batch of 32,768, and that one choice drives most of the GPU engineering on this page.

This article works from the loss function down to the hardware. It shows why the similarity matrix grows quadratically, how distributed training cuts that back to linear per GPU, what the memory and communication budget looks like for a real configuration, and which precision, data and accumulation mistakes quietly ruin runs.

Advertisement

The loss in one page

Take a batch of B image-caption pairs. An image tower, typically a Vision Transformer, and a text tower, a Transformer over tokens, each produce an embedding, and both are projected to the same dimension d and L2-normalised. The dot product of any image embedding with any caption embedding is then a cosine similarity. Stack them into a B by B matrix: entry (i, j) scores image i against caption j, and the diagonal holds the true pairs.

The loss is a cross-entropy over each row (for every image, pick its caption out of B candidates) and over each column (for every caption, pick its image), averaged. This is the symmetric InfoNCE objective. Before the softmax the similarities are multiplied by a learned scale, the inverse of a temperature. CLIP initialised the temperature at 0.07, so the scale starts near 14.3, and capped the scale at 100 to keep training stable.

Notice what B means here. Each image is scored against B minus 1 negatives. With 256 negatives the task is easy and the gradients carry little information. With 32,767 it is hard, and the model must learn fine distinctions. That is why contrastive models keep improving with batch sizes far beyond what a classifier would ever use, and why the GPU problem is mostly about making B large.

Why the logit matrix is the bottleneck

The towers' activations scale linearly with B, like any model. The similarity matrix scales with B squared. At B = 32,768 it holds about 1.07 billion entries, roughly 4.3 GB in fp32 for the logits alone, and the softmax and its gradient need about as much again. If every GPU materialised the full matrix, that cost would be paid on every device, on top of activations for a large ViT.

The CLIP paper already addressed this by sharding the similarity computation, so each GPU computes only the rows needed for its local embeddings. It also used mixed precision, gradient checkpointing and half-precision Adam statistics. The open-source reimplementation, open_clip, exposes the same idea as two flags, --local-loss and --gather-with-grad, and documents that they reduce the logit memory from quadratic to linear while producing numerically identical results.

Advertisement

Distributed local loss

Local loss with gather-with-grad: each GPU scores its rows against every columnGPU r: imageslocal batch bGPU r: captionslocal batch bImage towerViT, bf16 autocastText towertransformerL2 normaliseimg_r: b x dL2 normalisetxt_r: b x dall_gather (autograd)img_all, txt_all: (N*b) x dLocal logits (fp32)img_r @ txt_all.T and txt_r @ img_all.T : b x N*bCross-entropylabels = arange(b) + rank * bMemory per GPUnaive: (N*b)^2 logits local: b * N*bN = 256, b = 128: 1.07e9 vs 4.2e6 entriesgradients flow back to other ranks' towerslocal rows
Each rank computes embeddings for its own b pairs, gathers everyone's embeddings with an autograd-aware all-gather, and scores only its own rows against all N*b columns.

With N GPUs and a local batch of b, the global batch is N times b. Each rank runs both towers on its own pairs, then all-gathers the normalised embeddings from every rank. It now holds all N*b image and caption embeddings but computes logits only for its own b rows against all N*b columns, in both directions. Its positives sit at a fixed offset: local index i matches global column rank * b + i.

import torch
import torch.nn.functional as F
import torch.distributed as dist
import torch.distributed.nn as dist_nn          # differentiable collectives

def clip_loss(img_emb, txt_emb, logit_scale):
    # img_emb, txt_emb: (b, d) local, already L2-normalised; logit_scale: scalar param (log space)
    rank, world = dist.get_rank(), dist.get_world_size()
    img_all = torch.cat(dist_nn.all_gather(img_emb), dim=0)   # (world*b, d), grads flow back
    txt_all = torch.cat(dist_nn.all_gather(txt_emb), dim=0)
    scale = logit_scale.exp().clamp(max=100.0)
    logits_i = scale * img_emb.float() @ txt_all.float().T    # (b, world*b): my images vs all captions
    logits_t = scale * txt_emb.float() @ img_all.float().T    # my captions vs all images
    b = img_emb.shape[0]
    labels = torch.arange(b, device=img_emb.device) + rank * b  # my positives sit at my offset
    return 0.5 * (F.cross_entropy(logits_i, labels) + F.cross_entropy(logits_t, labels))

The all-gather must be differentiable. A plain dist.all_gather returns tensors that are detached from the graph, so the gradient of rank r's loss with respect to rank s's embeddings is silently dropped. Each tower then learns only from its own rows of the global loss, which weakens the training signal from negatives. The autograd version sends those gradients back to their owners during backward with a reduce-scatter, and that is exactly what --gather-with-grad enables. Getting this wrong does not crash. The loss still falls, just to a worse model, which is why it deserves a unit test against a single-GPU reference on a small batch.

Computing the logits and cross-entropy in fp32 costs little, because the local matrix is small, and avoids softmax precision problems at scales near 100.

Worked budget: 256 GPUs, global batch 32,768

QuantityNaive full matrixLocal loss
Logit entries per GPU, one direction32,768 squared, about 1.07e9128 x 32,768, about 4.2e6
fp32 logits, both directionsabout 8.6 GBabout 34 MB
Embeddings gathered per step (d = 768, bf16)32,768 x 768 x 2 B, about 50 MB per modalitysame
Gradient all-reduce (ViT-L/14 CLIP, about 430M params, bf16)about 0.86 GBsame

Two conclusions follow. First, local loss turns the logit matrix from the dominant memory cost into a rounding error, and the budget returns to tower activations, which you manage with gradient checkpointing and attention kernels. Second, the extra communication for contrastive training, about 100 MB of embeddings plus the matching reduce-scatter, is small next to the gradient all-reduce that any data-parallel job already pays. For how those collectives map onto NVLink and InfiniBand, see NCCL collectives.

For scale, the paper reports that its largest ResNet, RN50x64, trained for 18 days on 592 V100 GPUs and its largest ViT, ViT-L/14, for 12 days on 256 V100s, on 400 million image-text pairs. Modern GPUs with bf16, fused attention and better input pipelines cut this a long way, but the structure of the problem is unchanged.

The gradient accumulation trap

If you cannot fit a large local batch, the usual answer is gradient accumulation: run k micro-batches, sum the gradients, step once. For contrastive loss that is not equivalent to a k-times-larger batch. Each micro-batch computes its loss against only its own negatives, so you get k small-batch losses averaged together, and the model never sees the harder large-batch task.

The fix is to separate the embedding computation from the loss. First run all k micro-batches forward without gradients and cache the embeddings. Then compute the full contrastive loss over the cached set. Finally re-run each micro-batch forward with gradients, substituting its fresh embeddings into the cached set, and backpropagate. This costs one extra forward pass but gives the true large-batch gradient. The idea was published as GradCache, and open_clip implements the same approach behind --accum-freq, which the docs describe as simulating a batch of batch size times accumulation steps times GPU count.

torchrun --nproc_per_node 8 -m open_clip_train.main \
    --model ViT-B-16 --train-data '/data/shards/{00000..09999}.tar' --dataset-type webdataset \
    --train-num-samples 400000000 \
    --batch-size 256 --accum-freq 4 --precision amp_bf16 \
    --local-loss --gather-with-grad --grad-checkpointing \
    --lr 5e-4 --wd 0.2 --warmup 2000 --epochs 32 --workers 8

That command gives a global batch of 256 x 8 x 4 = 8,192 on one 8-GPU node. WebDataset training requires --train-num-samples; set it to your real sample count. Treat the hyperparameters as a starting point, not a recipe.

A minimal training loop

model = CLIP(vision="ViT-B-16", text_layers=12, embed_dim=512).cuda()
model = DDP(model, device_ids=[local_rank])                   # or FSDP for large towers
logit_scale = model.module.logit_scale                        # init log(1/0.07), kept in fp32
decay, no_decay = split_params(model)                         # no decay on norms, biases, logit_scale
opt = torch.optim.AdamW([{"params": decay, "weight_decay": 0.2},
                         {"params": no_decay, "weight_decay": 0.0}],
                        lr=5e-4, betas=(0.9, 0.98), eps=1e-6)

for step, (images, tokens) in enumerate(loader):              # WebDataset shards, distinct per rank
    images, tokens = images.cuda(non_blocking=True), tokens.cuda(non_blocking=True)
    lr_schedule(opt, step)                                    # warmup then cosine
    with torch.autocast("cuda", dtype=torch.bfloat16):
        img, txt = model(images, tokens)                      # call the DDP wrapper; forward returns normalised embeddings
    loss = clip_loss(img, txt, logit_scale)                   # logits in fp32
    opt.zero_grad(set_to_none=True)
    loss.backward()
    torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
    opt.step()
    if step % 100 == 0:
        log(loss=loss.item(), scale=logit_scale.exp().item(), samples_per_s=throughput())

A few details in this loop carry real weight. The logit scale is a parameter stored in log space, initialised to log(1/0.07), excluded from weight decay and clamped when used. Norm gains and biases are also excluded from decay. Encoders run under bf16 autocast while the loss runs in fp32. Gradient clipping guards against the occasional bad shard. Each rank must read different data shards, or the effective batch shrinks and duplicate pairs appear as false negatives.

Precision, kernels and memory

Use bf16 autocast on GPUs that support it. Its exponent range matches fp32, so the loss scaling that fp16 needs goes away, and --precision amp_bf16 selects it in open_clip. Keep the optimizer state and master weights in fp32 unless you have measured that lower-precision states are safe for your run. The trade-offs are covered in mixed precision training.

Tower activations dominate memory once local loss is on. Gradient checkpointing (--grad-checkpointing) recomputes block activations during backward and typically buys a much larger local batch for roughly a third more compute. Fused attention through PyTorch's scaled dot-product attention removes the sequence-squared attention matrix. FlashAttention explains why that is both faster and smaller. For towers too large to replicate on every GPU, shard parameters and optimizer state with FSDP. Contrastive loss works unchanged on top, because it only needs embeddings.

Patch dropout is a contrastive-specific accelerator. FLIP randomly drops a large share of image patches during pre-training (the paper studies 50 and 75 percent), so the image tower processes far fewer tokens per image and a larger batch fits in the same time. A short unmasked fine-tune at the end closes most of the gap.

Feeding the GPUs

CLIP runs are frequently input-bound rather than compute-bound. Hundreds of millions of small JPEGs cannot be read as individual files at speed, so datasets are packed into tar shards and streamed, the WebDataset format that open_clip supports. JPEG decoding and random-resized crops run on CPU workers. A node of eight fast GPUs training a ViT-B can consume many thousands of images per second, which takes dozens of decode cores. Watch GPU utilisation and data-loader wait time together. If utilisation drops whenever the loader queue empties, add workers, move decoding to the GPU, store images pre-resized near the training resolution, or lower resolution for most of training and raise it at the end.

Tokenisation is cheap but must be deterministic and truncated to the text tower's context length; CLIP used 77 tokens. Filtering matters more than speed. Duplicate and near-duplicate pairs inside a batch act as false negatives, and low-quality alt text wastes compute, which is why later open datasets invested so heavily in filtering.

SigLIP: a loss without the global softmax

SigLIP replaces the softmax with an independent sigmoid on every image-caption pair: positives are labelled 1 and negatives 0, with a learned temperature and bias, initialised to log 10 and minus 10 so that the flood of negatives does not swamp early training. Because no row needs normalising over all columns, the loss can be computed in chunks. Devices pass caption embeddings around a ring and accumulate pairwise terms without ever gathering the full set. The paper reports strong results at smaller batch sizes than softmax CLIP needs, which makes it attractive when your GPU count is modest. The engineering on this page still applies to the towers, data and precision.

Failure modes and diagnostics

SymptomLikely causeCheck
Loss stuck near ln(global batch)Embeddings collapsed or labels misalignedPositives at rank*b offset; gathered order matches ranks
Logit scale pinned at 100 earlyLearning rate too high or data too easy (duplicates)Lower LR; deduplicate shards
Multi-GPU worse than single-GPU at same global batchNon-differentiable gatherCompare gradients against a single-GPU reference
Loss spikes, then NaNfp16 overflow, bad shard, missing clipSwitch to bf16; log per-shard loss; clip grad norm
GPU utilisation saw-toothInput pipeline starvationLoader wait time; worker count; pre-resized shards
Zero-shot accuracy flat while loss fallsShards repeated across ranks or text truncation bugsPer-rank shard lists; inspect token lengths

Evaluate during training, not only at the end. Zero-shot ImageNet classification with prompt templates and image-text retrieval recall at 1 and 5 on a held-out set are cheap to run every few thousand steps and catch silent failures that the training loss hides.

What to do next

  1. Write the loss with a differentiable all-gather and unit-test it against a single-GPU computation on 2 ranks with a small batch.
  2. Turn on local loss, bf16 autocast, fp32 logits, logit-scale clamping and fused attention before tuning anything else.
  3. Find the largest local batch that fits with gradient checkpointing, then reach your target global batch with GPU count and feature-cached accumulation.
  4. Pack data into tar shards, assign distinct shards per rank, and measure loader wait time.
  5. Log loss, logit scale, gradient norm and throughput every few hundred steps; run zero-shot and retrieval evals on a schedule.
  6. If your GPU budget caps the batch well below 32K, run a SigLIP baseline alongside softmax CLIP.
  7. Consider patch dropout for long pre-training runs and finish with a short unmasked phase.
Key takeaway: In CLIP training the batch is the negative set, so batch size is a quality knob, and the B by B logit matrix is the first wall you hit. Distributed local loss with a differentiable all-gather makes the logit cost linear per GPU with identical results, leaving tower activations as the budget, managed with checkpointing, fused attention, bf16 and FSDP. Plain gradient accumulation silently shrinks the negative set; cache features and recompute instead. Keep the logit scale and loss in fp32, feed the GPUs from sharded data, test the gather against a single-GPU reference, and evaluate zero-shot during training.