Implementing Custom Data Loading Mechanisms for PyLate Training: A Complete Guide
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 inpylate/utils/processing.py): Transforms knowledge-distillation training sets by resolvingquery_idanddocument_idreferences to actual text strings from separatedatasets.Datasetobjects.ColBERTCollator(located inpylate/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.
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.
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):
# 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):
# 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.
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.
Complete End-to-End Implementation
The following snippet demonstrates the full pipeline for implementing custom data loading mechanisms for PyLate training:
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: Contains theKDProcessingclass that resolves IDs, truncates scores, and builds the final example dictionary. This is where then_waysfiltering logic resides. -
pylate/utils/collator.py: ImplementsColBERTCollator.__call__which detects columns likequery,documents, andscores, then tokenizes them using the model'stokenizemethod. See lines 98-124 for the column handling logic. -
pylate/losses/distillation.py: Defines the knowledge-distillation loss that consumes the data format produced byKDProcessing. It expects batched tensors containing query and document representations along with teacher scores. -
pylate/models/colbert.py: Provides thetokenizemethod used by the collator, handling prompt addition and tokenization specifics for ColBERT-style late interaction models. -
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
KDProcessingto join separate query and document datasets by resolving IDs to text on-the-fly, with configurable truncation via then_waysparameter. -
Choose between
map()andset_transform()depending on whether you need eager caching (map) or memory-efficient lazy loading (set_transform). -
Implement
ColBERTCollatorwith your model'stokenizefunction to convert text dictionaries into batched tensors suitable for PyLate's distillation losses. -
Reference the source files in
pylate/utils/processing.pyandpylate/utils/collator.pywhen 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 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.
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 →