# Implementing Knowledge Distillation Training Pipelines in PyLate: A Complete Guide

> Learn to implement knowledge distillation training pipelines in PyLate. Train a lightweight ColBERT student model using KL-divergence and soft labels from a teacher model. Get the full guide here.

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

---

**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`](https://github.com/lightonai/pylate/blob/main/pylate/losses/distillation.py). The `Distillation` class extends `torch.nn.Module` and implements the following logic:

- **Score Computation**: Uses `colbert_kd_scores` from `pylate/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.KLDivLoss` between log-softmax of student scores and teacher scores (labels), with configurable `size_average` parameter

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:

```python
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`](https://github.com/lightonai/pylate/blob/main/pylate/utils/processing.py) converts ID-based datasets into raw text suitable for tokenization. It performs three critical operations:

1. **Parsing**: Converts stringified Python literals (stored in dataset) back to lists using `ast.literal_eval`
2. **Truncation**: Limits documents to `n_ways` (default 32) to control memory usage and training time
3. **Resolution**: Maps query and document IDs to actual text strings using index maps

Apply the transformation using `set_transform`:

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

```python
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:

```python
from pylate import losses

# Initialize distillation loss

distill_loss = losses.Distillation(model=model)

```

Key implementation details from [`pylate/losses/distillation.py`](https://github.com/lightonai/pylate/blob/main/pylate/losses/distillation.py):
- **Score metric**: Defaults to `colbert_kd_scores` computing 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_softmax` on student predictions versus teacher labels, with reduction controlled by `size_average` parameter

## 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:

```python
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`](https://github.com/lightonai/pylate/blob/main/pylate/utils/processing.py) (lines 89-127) |
| **ColBERT** | Student architecture providing token-level embeddings and tokenization | [`pylate/models/colbert.py`](https://github.com/lightonai/pylate/blob/main/pylate/models/colbert.py) |
| **Distillation** | Computes KL-divergence between student similarities and teacher scores | [`pylate/losses/distillation.py`](https://github.com/lightonai/pylate/blob/main/pylate/losses/distillation.py) (lines 95-132) |
| **ColBERTCollator** | Generates skip-list masks for masked tokens during batching | [`pylate/utils/collator.py`](https://github.com/lightonai/pylate/blob/main/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`](https://github.com/lightonai/pylate/blob/main/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`](https://github.com/lightonai/pylate/blob/main/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 `Distillation` loss class in [`pylate/losses/distillation.py`](https://github.com/lightonai/pylate/blob/main/pylate/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 `KDProcessing` in [`pylate/utils/processing.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/processing.py), which handles truncation to `n_ways` (default 32) and text resolution.
- **Student architecture** uses `models.ColBERT` from [`pylate/models/colbert.py`](https://github.com/lightonai/pylate/blob/main/pylate/models/colbert.py), providing token-level embeddings and late-interaction scoring compatible with the distillation objective.
- **Training orchestration** delegates to `SentenceTransformerTrainer` from the `sentence-transformers` library, using `ColBERTCollator` for 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`](https://github.com/lightonai/pylate/blob/main/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`](https://github.com/lightonai/pylate/blob/main/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`](https://github.com/lightonai/pylate/blob/main/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.