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:

  1. Gathering tensors from all ranks using torch.distributed.all_gather.
  2. Reconstructing the full list while replacing the local rank’s position with the original tensor (preserving its gradient graph).
  3. 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:

  1. Launcher Setup: The user launches training with torchrun --nproc_per_node=4 python train.py (or equivalent), which sets WORLD_SIZE, RANK, MASTER_ADDR, and MASTER_PORT.

  2. 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.

  3. Transparent Utilities: Throughout training, get_rank() and get_world_size() provide context-aware values (returning 0 and 1 for single-GPU runs, actual values for distributed runs). The all_gather functions handle the communication when needed.

  4. Loss-Level Distribution: If CachedContrastive is configured with gather_across_devices=True, it automatically invokes all_gather_with_gradients during the forward pass, expanding the negative set across all GPUs while preserving gradients.

  5. 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.distributed through utility functions in pylate/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 like torchrun and initializes the NCCL backend.
  • The CachedContrastive loss enables cross-GPU negative mining through all_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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →