# How PyLate Manages and Supports Distributed Training Setups

> Discover how PyLate simplifies distributed training with torch distributed integration gradient preservation and seamless multi GPU scaling for efficient model development.

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

---

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

```python

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