Implementing Knowledge Distillation Training Pipelines in PyLate: A Complete Guide
Knowledge distillation in PyLate trains a lightweight ColBERT student model by minimizing KL-divergence between the student's similarity scores and soft labels generated by a stronger teacher model.
PyLate is an open-source framework for late-interaction retrieval models that simplifies implementing knowledge distillation (KD) pipelines. By leveraging the Distillation loss and KDProcessing utilities, you can transfer knowledge from large teacher models to efficient student architectures without manually implementing complex training loops.
Understanding Knowledge Distillation in PyLate
Knowledge distillation in PyLate follows the classic teacher-student paradigm adapted for late-interaction retrieval. The teacher model generates soft score distributions over candidate documents for each query, and the student learns to mimic these distributions rather than hard binary labels.
The Distillation Loss Architecture
The core implementation resides in pylate/losses/distillation.py. The Distillation class extends torch.nn.Module and implements the following logic:
- Score Computation: Uses
colbert_kd_scoresfrompylate/scores/__init__to compute dot-product similarities between query and document token embeddings - Normalization: Optional min-max normalization (lines 21-31) rescales teacher scores to
[0, 1]range, crucial when teacher models output raw dot-products with arbitrary scales - KL-Divergence: Computes
torch.nn.KLDivLossbetween log-softmax of student scores and teacher scores (labels), with configurablesize_averageparameter
The forward pass (lines 95-132) handles L2-normalized token embeddings, reshapes document tensors to (batch, n_ways, ...), generates skip-list masks via extract_skiplist_mask, and applies the scoring metric before the final divergence calculation.
Preparing Your Knowledge Distillation Dataset
PyLate requires specific dataset formatting for distillation workflows. Unlike standard supervised training, KD datasets contain pre-computed teacher scores alongside query and document identifiers.
Loading Teacher-Generated Data
Teacher-generated datasets typically contain three components:
- Queries: Mapping from query IDs to query text
- Documents: Mapping from document IDs to document text
- Training examples: Query ID, lists of document IDs, and corresponding teacher similarity scores
Load these using the Hugging Face datasets library:
from datasets import load_dataset
# Load teacher-generated training data
train = load_dataset("lightonai/ms-marco-en-bge", name="train")
# Load query and document corpora
queries = load_dataset("lightonai/ms-marco-en-bge", name="queries")
documents = load_dataset("lightonai/ms-marco-en-bge", name="documents")
Transforming IDs with KDProcessing
The KDProcessing class in pylate/utils/processing.py converts ID-based datasets into raw text suitable for tokenization. It performs three critical operations:
- Parsing: Converts stringified Python literals (stored in dataset) back to lists using
ast.literal_eval - Truncation: Limits documents to
n_ways(default 32) to control memory usage and training time - Resolution: Maps query and document IDs to actual text strings using index maps
Apply the transformation using set_transform:
from pylate import utils
# Initialize processor with query and document corpora
kd_processor = utils.KDProcessing(queries=queries, documents=documents)
# Attach transformation to dataset
train.set_transform(kd_processor.transform)
The transform method (lines 89-127) returns dictionaries containing query, documents, and scores fields, formatted for the Distillation loss.
Configuring the Student Model and Loss Function
PyLate implements the student architecture using the ColBERT class, which adds late-interaction capabilities to standard encoder backbones.
Initializing the ColBERT Student
The models.ColBERT class in pylate/models/colbert.py wraps any Hugging Face transformer, adding a linear projection layer when needed. It provides:
- Tokenization:
tokenize(is_query: bool)method producing token-level tensors with query/document specific handling - Embeddings: Forward pass returns
"token_embeddings"used by the distillation loss, L2-normalized per token
Initialize a lightweight student:
from pylate import models
import torch
# Create student from BERT-base
model = models.ColBERT(model_name_or_path="bert-base-uncased")
# Optional: Compile for speed (PyTorch 2.0+)
model = torch.compile(model)
Setting Up the Distillation Loss
The Distillation loss requires the student model instance and handles the complexity of late-interaction scoring:
from pylate import losses
# Initialize distillation loss
distill_loss = losses.Distillation(model=model)
Key implementation details from pylate/losses/distillation.py:
- Score metric: Defaults to
colbert_kd_scorescomputing MaxSim between query and document tokens - Normalization:
normalize_scores=True(default) applies min-max scaling to teacher scores, preventing gradient instability when teacher outputs have large magnitudes - KL-Divergence: Uses
log_softmaxon student predictions versus teacher labels, with reduction controlled bysize_averageparameter
Complete Training Pipeline Implementation
PyLate delegates training orchestration to the sentence-transformers library, leveraging SentenceTransformerTrainer for distributed training, mixed precision, and checkpointing.
The following script combines all components into a reproducible training pipeline:
import torch
from datasets import load_dataset
from sentence_transformers import (
SentenceTransformerTrainer,
SentenceTransformerTrainingArguments,
)
from pylate import losses, models, utils
# ----------------------------------------------------------------------
# 1️⃣ Load teacher-generated KD data
# ----------------------------------------------------------------------
train = load_dataset(
path="lightonai/ms-marco-en-bge",
name="train",
)
queries = load_dataset(
path="lightonai/ms-marco-en-bge",
name="queries",
)
documents = load_dataset(
path="lightonai/ms-marco-en-bge",
name="documents",
)
# ----------------------------------------------------------------------
# 2️⃣ Attach processing that resolves IDs → text & truncates scores
# ----------------------------------------------------------------------
train.set_transform(
utils.KDProcessing(queries=queries, documents=documents).transform,
)
# ----------------------------------------------------------------------
# 3️⃣ Define student ColBERT model
# ----------------------------------------------------------------------
model = models.ColBERT(model_name_or_path="bert-base-uncased")
model = torch.compile(model) # optional speed-up on supported hardware
# ----------------------------------------------------------------------
# 4️⃣ Prepare trainer args
# ----------------------------------------------------------------------
run_name = "knowledge-distillation-bert-base"
output_dir = f"output/{run_name}"
args = SentenceTransformerTrainingArguments(
output_dir=output_dir,
num_train_epochs=1,
per_device_train_batch_size=16,
fp16=True, # set False if GPU lacks FP16 support
run_name=run_name,
learning_rate=1e-5,
)
# ----------------------------------------------------------------------
# 5️⃣ Distillation loss
# ----------------------------------------------------------------------
distill_loss = losses.Distillation(model=model)
# ----------------------------------------------------------------------
# 6️⃣ Trainer – note the collator aligns with skip-list masking
# ----------------------------------------------------------------------
trainer = SentenceTransformerTrainer(
model=model,
args=args,
train_dataset=train,
loss=distill_loss,
data_collator=utils.ColBERTCollator(tokenize_fn=model.tokenize),
)
# ----------------------------------------------------------------------
# 7️⃣ Run training
# ----------------------------------------------------------------------
trainer.train()
Explanation of key components
| Component | Purpose | Source Reference |
|---|---|---|
| KDProcessing | Converts dataset IDs to raw texts and truncates to n_ways (default 32) |
pylate/utils/processing.py (lines 89-127) |
| ColBERT | Student architecture providing token-level embeddings and tokenization | pylate/models/colbert.py |
| Distillation | Computes KL-divergence between student similarities and teacher scores | pylate/losses/distillation.py (lines 95-132) |
| ColBERTCollator | Generates skip-list masks for masked tokens during batching | pylate/utils/collator.py |
| SentenceTransformerTrainer | Handles distributed training, mixed precision, and optimization | sentence-transformers library |
Key Implementation Details and Optimization
Handling Skip-List Masks and Tokenization
The ColBERTCollator (referenced in pylate/utils/collator.py) works with the Distillation loss to handle skip-list masks. These masks identify tokens that should be ignored during similarity computation (e.g., punctuation or special tokens). The collator generates these masks during batching, and the loss applies them via extract_skiplist_mask before computing the MaxSim operation in colbert_kd_scores.
Score Normalization Strategies
When implementing knowledge distillation training pipelines in PyLate, teacher score normalization is critical for training stability. The Distillation class in pylate/losses/distillation.py provides normalize_scores=True by default, which applies min-max scaling (lines 21-31) to rescale teacher scores to the [0, 1] range. This prevents gradient explosion when teachers output raw dot-products with large magnitudes. Disable this only if your teacher already outputs calibrated probabilities.
Distributed Training Considerations
The Distillation loss is compatible with PyTorch Distributed Data Parallel (DDP) because it gracefully handles wrapped models. When accessing model attributes like skiplist and do_query_expansion, the loss checks both the model and model.module (lines 54-68), ensuring seamless operation whether training on a single GPU or across multiple nodes.
Summary
- PyLate implements knowledge distillation through the
Distillationloss class inpylate/losses/distillation.py, which computes KL-divergence between student similarity scores and teacher-generated soft labels. - Dataset preparation requires converting ID-based datasets to raw text using
KDProcessinginpylate/utils/processing.py, which handles truncation ton_ways(default 32) and text resolution. - Student architecture uses
models.ColBERTfrompylate/models/colbert.py, providing token-level embeddings and late-interaction scoring compatible with the distillation objective. - Training orchestration delegates to
SentenceTransformerTrainerfrom thesentence-transformerslibrary, usingColBERTCollatorfor skip-list mask generation and supporting distributed training with automatic model unwrapping.
Frequently Asked Questions
How does the Distillation loss handle different teacher score ranges?
The Distillation class automatically normalizes teacher scores to the [0, 1] range using min-max scaling when normalize_scores=True (the default). This occurs in pylate/losses/distillation.py (lines 21-31) and prevents training instability when teachers output unbounded dot-product similarities. You can disable this if your teacher already produces calibrated probabilities between 0 and 1.
What is the purpose of KDProcessing in the training pipeline?
KDProcessing in pylate/utils/processing.py bridges the gap between ID-based datasets and the text inputs required by ColBERT models. It converts stringified document ID lists and teacher scores into raw texts using ast.literal_eval, truncates examples to n_ways (default 32) to control memory usage, and returns dictionaries with query, documents, and scores fields ready for the Distillation loss.
Can I use knowledge distillation with multiple GPUs or distributed training?
Yes, the PyLate distillation pipeline supports Distributed Data Parallel (DDP) through the SentenceTransformerTrainer. The Distillation loss in pylate/losses/distillation.py specifically handles DDP model wrapping by checking for attributes on both model and model.module (lines 54-68), ensuring skip-list masks and query expansion settings are accessible regardless of whether the model is wrapped for distributed training.
How do I customize the similarity scoring function in the distillation loss?
The Distillation class accepts a score_metric parameter that defaults to colbert_kd_scores from pylate/scores/__init__. You can substitute this with alternative scoring functions (such as cosine similarity or inner-product variants) by passing a callable that accepts query embeddings, document embeddings, and skip-list masks, then returns similarity scores. This allows experimentation with different late-interaction scoring mechanisms while maintaining the same KL-divergence training objective.
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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →