How PyLate Manages and Supports Distributed Training Setups
PyLate abstracts distributed training through a lightweight wrapper around torch.distributed, providing gradient-preserving all-gather operations and transparent multi-GPU scaling via the CachedContrastive loss with gather_across_devices=True.
PyLate, an open-source library maintained by lightonai/pylate, simplifies training and inference for late-interaction retrieval models like ColBERT. When scaling to multiple GPUs or nodes, the library handles the complexity of distributed training automatically, allowing the same training script to run unchanged on a single GPU or a multi-node cluster.
Core Distributed Training Components in PyLate
PyLate’s distributed architecture centers on two utility modules that wrap PyTorch’s native distributed backend, plus a specialized loss function that leverages these utilities for cross-GPU negative mining.
Utility Functions in pylate.utils.distributed
The primary interface for distributed operations resides in pylate/utils/distributed.py. This module exposes four critical functions that automatically detect whether distributed training is active and fall back to no-op behavior (with a single warning) when running on a single device:
all_gather(tensor): Gathers tensors from all ranks into a list, preserving the autograd graph only for the local rank’s tensor.all_gather_with_gradients(tensor): Gathers tensors while preserving gradients for all ranks, enabling backpropagation through embeddings collected from other GPUs.get_rank(): Returns the current process rank (0 if not distributed).get_world_size(): Returns the total number of processes (1 if not distributed).
Process Group Initialization
For multi-node setups, pylate/indexes/stanford_nlp/utils/distributed.py handles the low-level process group creation. The init(rank) function reads WORLD_SIZE, RANK, MASTER_ADDR, and MASTER_PORT from the environment (typically set by torchrun or mpirun), initializes the NCCL backend via torch.distributed.init_process_group, and binds each rank to its corresponding CUDA device.
This module also provides a barrier() helper that synchronizes all workers after I/O-heavy operations like index building.
Gradient-Preserving All-Gather Operations
The key innovation in PyLate’s distributed support is the gradient-preserving gather mechanism. Standard torch.distributed.all_gather breaks the autograd graph because it creates new tensors without history. PyLate’s all_gather_with_gradients solves this by:
- Gathering tensors from all ranks using
torch.distributed.all_gather. - Reconstructing the full list while replacing the local rank’s position with the original tensor (preserving its gradient graph).
- For non-local ranks, the gathered tensors participate in the loss computation but do not require gradients back to their origin ranks (the local rank only backprops through its own embeddings).
This allows the CachedContrastive loss to mine negatives from all GPUs while maintaining correct gradient flow.
Distributed Negative Mining with CachedContrastive
PyLate enables large-scale contrastive learning through the CachedContrastive loss in pylate/losses/cached_contrastive.py. When initialized with gather_across_devices=True, the loss automatically expands the negative pool to include documents from all GPUs.
During the forward pass (lines 49-57), the loss performs:
# Gather document embeddings from all ranks with gradient preservation
embeddings_other = [torch.cat(all_gather_with_gradients(emb)) for emb in embeddings_other]
# Gather masks (no gradients needed)
masks = [masks[0], *[torch.cat(all_gather(m)) for m in masks[1:]]]
# Adjust labels to account for gathered batches
labels = labels + get_rank() * batch_size
This implementation multiplies the effective batch size by world_size without increasing per-GPU memory usage, as only the local embeddings are cached locally while the gathered tensors are used transiently for the contrastive computation.
How the Distributed Components Work Together
PyLate’s distributed training follows a seamless integration pattern that requires no boilerplate code in user scripts:
-
Launcher Setup: The user launches training with
torchrun --nproc_per_node=4 python train.py(or equivalent), which setsWORLD_SIZE,RANK,MASTER_ADDR, andMASTER_PORT. -
Automatic Initialization: When the training script imports PyLate and creates a model, the underlying utilities automatically detect the distributed environment. The
stanford_nlp.utils.distributed.init()function establishes the NCCL process group and assigns the correct CUDA device to each rank. -
Transparent Utilities: Throughout training,
get_rank()andget_world_size()provide context-aware values (returning 0 and 1 for single-GPU runs, actual values for distributed runs). Theall_gatherfunctions handle the communication when needed. -
Loss-Level Distribution: If
CachedContrastiveis configured withgather_across_devices=True, it automatically invokesall_gather_with_gradientsduring the forward pass, expanding the negative set across all GPUs while preserving gradients. -
Synchronization: After I/O operations like index serialization,
barrier()ensures all ranks synchronize before proceeding.
This architecture means a single training script works unchanged across single-GPU, multi-GPU single-node, and multi-node configurations.
Key Implementation Files
| File | Purpose |
|---|---|
pylate/utils/distributed.py |
Core wrappers around torch.distributed providing all_gather, all_gather_with_gradients, get_rank, and get_world_size with automatic fallback for single-device training. |
pylate/indexes/stanford_nlp/utils/distributed.py |
Process group initialization using NCCL backend and environment variable detection (WORLD_SIZE, RANK, MASTER_ADDR, MASTER_PORT), plus barrier synchronization. |
pylate/losses/cached_contrastive.py |
Implements CachedContrastive loss with optional gather_across_devices parameter that leverages all_gather_with_gradients for cross-GPU negative mining. |
pylate/models/colbert.py |
ColBERT model definition that operates transparently under distributed data parallel; relies on utility functions for any distributed-aware operations. |
examples/train/reason_moderncolbert.py |
Reference training script demonstrating distributed training configuration with gather_across_devices=True and no explicit distributed boilerplate. |
Summary
- PyLate abstracts
torch.distributedthrough utility functions inpylate/utils/distributed.py, providing gradient-preserving gathers and rank-aware helpers that fallback gracefully to single-GPU mode. - Process group initialization is handled automatically via
pylate/indexes/stanford_nlp/utils/distributed.py, which detects environment variables set by launchers liketorchrunand initializes the NCCL backend. - The
CachedContrastiveloss enables cross-GPU negative mining throughall_gather_with_gradients, effectively multiplying batch size across devices without breaking the autograd graph. - Training scripts require no distributed boilerplate; the same code runs on single GPU, multi-GPU single node, or multi-node clusters when launched with appropriate environment variables.
Frequently Asked Questions
Does PyLate require code changes to run on multiple GPUs?
No. PyLate training scripts work unchanged across hardware configurations. When launched with torchrun or similar distributed launchers, the library automatically detects the distributed environment through WORLD_SIZE and RANK environment variables and initializes the NCCL process group. The utility functions in pylate/utils/distributed.py provide fallback behavior for single-GPU runs, ensuring the same code path executes regardless of the number of available GPUs.
How does PyLate preserve gradients during distributed gathering?
PyLate implements all_gather_with_gradients in pylate/utils/distributed.py to maintain the autograd graph across ranks. Unlike standard torch.distributed.all_gather, which creates new tensors without gradient history, PyLate's implementation gathers tensors from all ranks but replaces the local rank's position in the gathered list with the original tensor. This preserves the local gradient graph while allowing the loss computation to utilize embeddings from other GPUs. The non-local tensors participate in forward computation but do not require gradient propagation back to their origin ranks.
What launchers are compatible with PyLate distributed training?
PyLate is compatible with any launcher that sets the standard PyTorch distributed environment variables: WORLD_SIZE, RANK, LOCAL_RANK, MASTER_ADDR, and MASTER_PORT. This includes torchrun (formerly torch.distributed.launch), mpirun with appropriate wrappers, Kubernetes PyTorch operators, and SLURM srun when configured for PyTorch distributed. The pylate/indexes/stanford_nlp/utils/distributed.py module reads these variables to initialize the NCCL backend via torch.distributed.init_process_group.
Can PyLate run on CPU-only environments?
Yes, though with limitations. The utility functions in pylate/utils/distributed.py automatically detect when torch.distributed is not initialized and return fallback values (rank=0, world_size=1) while emitting a single warning. This allows single-CPU training. However, the process group initialization in pylate/indexes/stanford_nlp/utils/distributed.py specifically uses the NCCL backend, which requires CUDA. For CPU-only distributed training, users would need to modify the backend to Gloo, though this is not the primary use case for PyLate, which is optimized for GPU-accelerated late interaction retrieval.
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 →