Differentiating PyLate's Contrastive, Distillation, and CachedContrastive Losses: A Complete Guide

PyLate provides three distinct loss functions—Contrastive, Distillation, and CachedContrastive—each optimized for different training scenarios ranging from standard contrastive learning to knowledge distillation and memory-constrained large-batch training.

Differentiating PyLate's Contrastive, Distillation, and CachedContrastive losses is essential for optimizing ColBERT-style dense retrieval models in the lightonai/pylate repository. Each loss implements a unique optimization strategy: standard in-batch negative sampling, teacher-student knowledge transfer, or gradient-cached massive batch training. Understanding these architectural differences ensures you select the right loss for your hardware constraints and training objectives.

Contrastive Loss: Standard In-Batch Negatives

The Contrastive loss implements the foundational ColBERT training objective, directly optimizing query-document relevance through in-batch negative sampling.

How It Works

In pylate/losses/contrastive.py, the forward method (lines 31-114) processes all sentence features in a single pass. It computes normalized token embeddings for every query and document, constructs a similarity matrix using colbert_scores (max-sim operator), and applies cross-entropy loss over the scores. The loss treats each query's positive document as the target class and all other documents in the batch as negatives.

Memory Characteristics and Multi-GPU Support

The standard Contrastive loss holds the complete forward computation graph in memory, making it suitable for small-to-moderate batch sizes. It supports gather_across_devices, which uses all_gather_with_gradients to pool document embeddings across GPUs, effectively multiplying the negative pool size without increasing per-GPU memory usage.

from pylate import models, losses

model = models.ColBERT(
    model_name_or_path="sentence-transformers/all-MiniLM-L6-v2",
    device="cpu",
)

# Standard contrastive loss

contrastive = losses.Contrastive(model=model, gather_across_devices=False)

anchor = model.tokenize(["fruits are healthy."], is_query=True)
positive = model.tokenize(["fruits are good for health."], is_query=False)
negative = model.tokenize(["fruits are bad for health."], is_query=False)

sentence_features = [anchor, positive, negative]
loss = contrastive(sentence_features=sentence_features)
print("Contrastive loss:", loss.item())

Distillation Loss: Knowledge Transfer from Teacher Models

The Distillation loss enables training smaller student models using soft targets from larger teacher models, diverging completely from the contrastive paradigm.

Teacher-Student Architecture

Implemented in pylate/losses/distillation.py, this loss uses colbert_kd_scores instead of standard colbert_scores. The forward method (lines 34-69) generates query and document embeddings, reshapes documents to (batch, n_ways, ...) to handle multiple candidates per query, and computes similarity scores between the student model's representations.

KL-Divergence Implementation

Rather than cross-entropy against hard labels, Distillation computes KL-divergence using torch.nn.KLDivLoss between the student's log-softmax scores and the teacher's soft targets. The loss optionally normalizes scores to the [0,1] range before computing divergence, ensuring stable gradient flow when transferring knowledge from high-capacity teachers.

import torch
from pylate import models, losses

model = models.ColBERT(
    model_name_or_path="sentence-transformers/all-MiniLM-L6-v2",
    device="cpu",
)

# Distillation loss for teacher-student training

distillation = losses.Distillation(model=model)

query = model.tokenize(["fruits are healthy."], is_query=True)
documents = model.tokenize(
    ["fruits are good for health.", "fruits are bad for health."],
    is_query=False,
)

sentence_features = [query, documents]

# Teacher logits (e.g., from a larger model); must sum to 1 per query

teacher_logits = torch.tensor([[0.7, 0.3]], dtype=torch.float32)

loss = distillation(sentence_features=sentence_features, labels=teacher_logits)
print("Distillation loss:", loss.item())

CachedContrastive Loss: Scaling to Massive Batch Sizes

The CachedContrastive loss solves the memory bottleneck of standard contrastive training, enabling effective batch sizes orders of magnitude larger than GPU memory would normally allow.

Gradient Caching Mechanism

Inspired by GradCache, this loss (implemented in pylate/losses/cached_contrastive.py) avoids storing the complete forward computation graph. Instead, it employs a two-pass strategy:

  1. Caching Pass: Runs with torch.no_grad(), computing and storing embeddings and random states needed for exact reproduction.
  2. Gradient Pass: Recomputes embeddings with gradients enabled, attaches a backward hook (_backward_hook), and injects previously cached gradients during backpropagation.

This "gradient-checkpointing" style approach trades computation for memory, allowing the loss to scale to thousands of in-batch negatives.

Mini-Batch Processing

The embed_minibatch_iter function (lines 83-114) splits sentence features into chunks specified by mini_batch_size. The calculate_loss_and_cache_gradients method processes these chunks, accumulating gradients in self.cache before the final cross-entropy computation over the full-batch score matrix.

from pylate import models, losses

model = models.ColBERT(
    model_name_or_path="sentence-transformers/all-MiniLM-L6-v2",
    device="cpu",
)

# CachedContrastive for large-batch training

cached = losses.CachedContrastive(
    model=model,
    mini_batch_size=2,          # small chunks to fit memory

    gather_across_devices=False,
    show_progress_bar=False,
)

anchors = model.tokenize(
    ["fruits are healthy.", "chips are not healthy."],
    is_query=True,
)
positives = model.tokenize(
    ["fruits are good for health.", "chips are not good for health."],
    is_query=False,
)
negatives = model.tokenize(
    ["fruits are bad for health.", "chips are good for health."],
    is_query=False,
)

sentence_features = [anchors, positives, negatives]
loss = cached(sentence_features=sentence_features)
print("CachedContrastive loss:", loss.item())

Key Differences and Selection Guide

Choosing between these losses depends on your training infrastructure, data availability, and model architecture:

Loss Optimization Target Memory Profile Best For
Contrastive In-batch negative contrast High (full graph) Standard training with moderate batch sizes
Distillation Teacher-student KL divergence Medium (no cross-batch negatives) Model compression, transferring from large teachers
CachedContrastive In-batch negatives with gradient caching Low (mini-batch chunks) Large-scale training with thousands of negatives

Contrastive is the default choice for most ColBERT training scenarios. Use Distillation when you have access to a high-quality teacher model and want to train a smaller, faster student. Deploy CachedContrastive when you need to maximize in-batch negatives for hard negative mining but face GPU memory constraints.

Summary

  • Contrastive implements standard in-batch negative training with full forward graphs, suitable for moderate batch sizes and multi-GPU setups with gather_across_devices.
  • Distillation transfers knowledge from teacher models using KL-divergence over colbert_kd_scores, ideal for model compression scenarios.
  • CachedContrastive enables massive effective batch sizes through gradient caching and mini-batch processing, trading computation for memory efficiency.
  • All three losses operate on ColBERT-style token embeddings and utilize colbert_scores or colbert_kd_scores from pylate/scores/similarity_functions.py.

Frequently Asked Questions

What is the main memory advantage of CachedContrastive over standard Contrastive loss?

CachedContrastive reduces memory usage by splitting the batch into mini-batches and employing gradient caching. Instead of storing the complete forward computation graph for all samples, it caches embeddings and gradients during a no-grad pass, then recomputes embeddings with a backward hook that injects cached gradients. This allows training with thousands of in-batch negatives that would otherwise cause out-of-memory errors in the standard Contrastive loss.

When should I use Distillation instead of Contrastive loss?

Use Distillation when you have access to a larger, pre-trained teacher model and want to transfer its knowledge to a smaller student ColBERT model. Unlike Contrastive, which learns from hard in-batch negatives, Distillation optimizes the student to match the teacher's soft target distributions using KL-divergence. This is particularly effective for model compression or when the teacher provides high-quality relevance scores that are expensive to compute at inference time.

How does gather_across_devices work in PyLate losses?

The gather_across_devices parameter enables multi-GPU training by gathering document embeddings across all processes using all_gather_with_gradients. In both Contrastive and CachedContrastive, setting this to True effectively multiplies your negative pool size by the number of GPUs, improving hard negative mining without increasing per-GPU memory usage. The gathered tensors maintain gradient flow, allowing the loss to backpropagate through the full distributed batch.

Can I switch between Contrastive and CachedContrastive without changing my data pipeline?

Yes, both losses share the same input interface and expect identical sentence_features formatting (lists of tokenized query and document batches). You can swap Contrastive for CachedContrastive by simply changing the loss class and specifying a mini_batch_size that fits your GPU memory. No changes to your data loading or tokenization logic are required, though you may want to increase your effective batch size when using CachedContrastive to take advantage of its memory efficiency.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →