# Implementing Custom Data Loading Mechanisms for PyLate Training: A Complete Guide

> Learn to implement custom data loading for PyLate training using KDProcessing and ColBERTCollator. Efficiently batch training data for knowledge distillation workflows.

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

---

**Use `KDProcessing` and `ColBERTCollator` from the PyLate utilities to resolve query/document IDs on-the-fly and batch training data for knowledge distillation workflows.**

PyLate is a neural search library built on top of Sentence-Transformers that requires specialized data handling for ColBERT-style training. Implementing custom data loading mechanisms for PyLate training involves leveraging the library's `KDProcessing` class to join separate query and document datasets, and the `ColBERTCollator` to prepare batched tensors for the trainer.

## Understanding PyLate's Data Loading Architecture

PyLate relies on the 🤗 *datasets* library for feeding training data, but introduces two critical abstractions for custom loading:

- **`KDProcessing`** (located in [`pylate/utils/processing.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/processing.py)): Transforms knowledge-distillation training sets by resolving `query_id` and `document_id` references to actual text strings from separate `datasets.Dataset` objects.
- **`ColBERTCollator`** (located in [`pylate/utils/collator.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/collator.py)): Converts raw feature dictionaries into batched tensor dictionaries that the PyLate model consumes during training.

Both utilities are stateless after construction, making them safe to pass to `datasets.Dataset.set_transform` or `SentenceTransformerTrainer` via the `data_collator` argument.

## Loading Separate Query and Document Datasets

Knowledge distillation in PyLate typically requires three distinct dataset splits: training examples (containing IDs and teacher scores), a query lookup table, and a document lookup table.

```python
from datasets import load_dataset

# Load the three splits that constitute a knowledge-distillation training set

train = load_dataset("lightonai/ms-marco-en-bge", name="train")          # scores + ids

queries = load_dataset("lightonai/ms-marco-en-bge", name="queries")      # id → text

documents = load_dataset("lightonai/ms-marco-en-bge", name="documents")  # id → text

```

The training split contains `query_id` and `document_id` fields that reference entries in the `queries` and `documents` datasets respectively.

## Implementing Knowledge Distillation Data Processing

### Configuring KDProcessing

The `KDProcessing` class handles the resolution of IDs to text and manages the truncation of teacher scores to a manageable number of negatives per query.

```python
from pylate.utils import KDProcessing

# Create a processor that resolves ids → raw texts at read time

processor = KDProcessing(queries=queries, documents=documents, n_ways=32)

```

The `n_ways` parameter (defaulting to 32) limits how many teacher scores and documents are retained per training example, controlling memory usage and training batch composition.

### Mapping vs. Lazy Transformations

PyLate supports two patterns for applying `KDProcessing` to your dataset:

**Eager mapping** (materializes a new dataset with resolved texts):

```python

# Map-style processing creates a new dataset with query and documents fields

train = train.map(processor.map)

```

**Lazy transformation** (resolves IDs on-the-fly during iteration):

```python

# set_transform keeps the original dataset immutable and performs lookup only when accessed

train.set_transform(processor.transform)

```

Use `map()` when you want to cache the resolved texts to disk, and `set_transform()` when working with large document collections that cannot fit in memory or when you need dynamic sampling.

## Batching with ColBERTCollator

The `ColBERTCollator` bridges the gap between resolved text examples and the tensor batches required by PyLate's `Distillation` loss.

```python
from pylate.utils import ColBERTCollator
from pylate import models

model = models.ColBERT(model_name_or_path="bert-base-uncased")
collator = ColBERTCollator(tokenize_fn=model.tokenize)

```

The collator automatically detects columns named `query`, `positive`, `negative`, `documents`, and others, applying the model's `tokenize` method (including any configured prompts) to produce input IDs, attention masks, and token type IDs.

When integrated with `SentenceTransformerTrainer`, the collator ensures that variable-length document lists are properly padded and batched for the knowledge distillation loss implemented in [`pylate/losses/distillation.py`](https://github.com/lightonai/pylate/blob/main/pylate/losses/distillation.py).

## Complete End-to-End Implementation

The following snippet demonstrates the full pipeline for implementing custom data loading mechanisms for PyLate training:

```python
from datasets import load_dataset
from pylate.utils import KDProcessing, ColBERTCollator
from pylate import models, losses
from sentence_transformers import SentenceTransformerTrainer, SentenceTransformerTrainingArguments

# ── Load raw datasets ───────────────────────────────────────────────────────

train = load_dataset("lightonai/ms-marco-en-bge", name="train")
queries = load_dataset("lightonai/ms-marco-en-bge", name="queries")
documents = load_dataset("lightonai/ms-marco-en-bge", name="documents")

# ── Prepare processing ─────────────────────────────────────────────────────

processor = KDProcessing(queries=queries, documents=documents, n_ways=32)
train = train.map(processor.map)                     # materialise enriched training set

# ── Model & collator ───────────────────────────────────────────────────────

model = models.ColBERT(model_name_or_path="bert-base-uncased")
collator = ColBERTCollator(tokenize_fn=model.tokenize)

# ── Trainer config ─────────────────────────────────────────────────────────

args = SentenceTransformerTrainingArguments(
    output_dir="output/kd-bert-base",
    num_train_epochs=3,
    per_device_train_batch_size=16,
    fp16=True,
)

trainer = SentenceTransformerTrainer(
    model=model,
    args=args,
    train_dataset=train,
    loss=losses.Distillation(model=model),
    data_collator=collator,
)

trainer.train()

```

This implementation automatically resolves query and document IDs to raw text, truncates teacher scores to the top-32 documents per query, and batches the data for efficient ColBERT training.

## Key Source Files and Implementation Details

Understanding the underlying source code helps when debugging or extending PyLate's data loading:

- **[`pylate/utils/processing.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/processing.py)**: Contains the `KDProcessing` class that resolves IDs, truncates scores, and builds the final example dictionary. This is where the `n_ways` filtering logic resides.

- **[`pylate/utils/collator.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/collator.py)**: Implements `ColBERTCollator.__call__` which detects columns like `query`, `documents`, and `scores`, then tokenizes them using the model's `tokenize` method. See lines 98-124 for the column handling logic.

- **[`pylate/losses/distillation.py`](https://github.com/lightonai/pylate/blob/main/pylate/losses/distillation.py)**: Defines the knowledge-distillation loss that consumes the data format produced by `KDProcessing`. It expects batched tensors containing query and document representations along with teacher scores.

- **[`pylate/models/colbert.py`](https://github.com/lightonai/pylate/blob/main/pylate/models/colbert.py)**: Provides the `tokenize` method used by the collator, handling prompt addition and tokenization specifics for ColBERT-style late interaction models.

- **[`examples/train/knowledge_distillation.py`](https://github.com/lightonai/pylate/blob/main/examples/train/knowledge_distillation.py)**: A complete working example in the repository that mirrors the implementation patterns described above.

These files together define the customizable data-loading pipeline that PyLate expects for both contrastive and knowledge-distillation training.

## Summary

- **Use `KDProcessing`** to join separate query and document datasets by resolving IDs to text on-the-fly, with configurable truncation via the `n_ways` parameter.

- **Choose between `map()` and `set_transform()`** depending on whether you need eager caching (`map`) or memory-efficient lazy loading (`set_transform`).

- **Implement `ColBERTCollator`** with your model's `tokenize` function to convert text dictionaries into batched tensors suitable for PyLate's distillation losses.

- **Reference the source files** in [`pylate/utils/processing.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/processing.py) and [`pylate/utils/collator.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/collator.py) when extending or debugging the data pipeline.

## Frequently Asked Questions

### How does KDProcessing handle large document collections that don't fit in memory?

`KDProcessing` stores references to the `queries` and `documents` datasets rather than copying them. When using `set_transform(processor.transform)`, text resolution occurs lazily during iteration, keeping memory usage constant regardless of corpus size. For extremely large collections, ensure your document dataset is memory-mapped or streamed using 🤗 Datasets' `streaming=True` mode.

### Can I use custom column names with ColBERTCollator?

The `ColBERTCollator` automatically detects standard columns including `query`, `positive`, `negative`, `documents`, and `scores` as implemented in [`pylate/utils/collator.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/collator.py) lines 98-124. If your dataset uses different column names, either rename them before collation or subclass `ColBERTCollator` and override the `__call__` method to map your custom fields to the expected keys.

### What is the optimal value for the n_ways parameter in KDProcessing?

The `n_ways` parameter defaults to 32, which balances training efficiency with memory usage for most knowledge distillation scenarios. Higher values (64-128) provide more negatives per query but increase GPU memory consumption during training. Lower values (8-16) reduce memory pressure but may limit the effectiveness of the distillation loss. Adjust based on your available VRAM and the complexity of your teacher model's score distribution.

### How do I debug issues with my custom data loading pipeline?

Start by verifying that `KDProcessing` correctly resolves IDs by inspecting a single example after applying `processor.map` or `processor.transform`. Check that `query` and `documents` fields contain actual strings rather than IDs. Next, test the `ColBERTCollator` independently by passing a list of examples to `collator(examples)` and verifying the output contains `input_ids`, `attention_mask`, and other required tensors. Finally, enable `logging_steps=1` in your `SentenceTransformerTrainingArguments` to trace data flow during the first training steps.