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

> Explore PyLates Contrastive Distillation and CachedContrastive losses Understand their unique strengths for contrastive learning knowledge distillation and efficient large batch training.

- Repository: [LightOn/pylate](https://github.com/lightonai/pylate)
- Tags: deep-dive
- Published: 2026-03-06

---

**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`](https://github.com/lightonai/pylate/blob/main/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.

```python
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`](https://github.com/lightonai/pylate/blob/main/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.

```python
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`](https://github.com/lightonai/pylate/blob/main/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.

```python
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`](https://github.com/lightonai/pylate/blob/main/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.