# How to Configure Multi-GPU Training in PyLate for ColBERT Models

> Master multi-GPU training for ColBERT models in PyLate. Configure gather across devices and launch with torchrun for efficient distributed training across multiple GPUs.

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

---

**Enable multi-GPU training in PyLate by setting `gather_across_devices=True` in your loss function and launching with `torchrun --nproc_per_node=N` to distribute ColBERT training across multiple GPUs.**

PyLate is an open-source library by LightOn AI that implements ColBERT-style late-interaction retrieval models. Configuring multi-GPU training in PyLate allows you to scale effective batch sizes and reduce training time by leveraging `torch.distributed` across multiple CUDA devices.

## Prerequisites and Installation

Install PyLate with the development dependencies and PyTorch:

```bash
pip install "pylate[dev]" torch torchvision

```

Optionally install Hugging Face Accelerate for simplified multi-GPU launching:

```bash
pip install accelerate

```

PyLate’s core distributed utilities are located in [`pylate/utils/distributed.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/distributed.py), which provides wrappers for process group initialization and cross-device tensor gathering.

## Setting Up the Training Script for Multi-GPU

The key to multi-GPU training lies in configuring your loss function to gather embeddings across all devices. The [`examples/train/contrastive.py`](https://github.com/lightonai/pylate/blob/main/examples/train/contrastive.py) file demonstrates the complete setup.

### Essential Code Components

```python

# examples/train/contrastive.py

from sentence_transformers import SentenceTransformerTrainer, SentenceTransformerTrainingArguments
from pylate import models, losses, utils, evaluation
import torch

# Initialize ColBERT model

model = models.ColBERT(model_name_or_path="bert-base-uncased")
model = torch.compile(model)  # Optional compilation for speed

# Load dataset

from datasets import load_dataset
dataset = load_dataset("sentence-transformers/msmarco-bm25", "triplet", split="train")
splits = dataset.train_test_split(test_size=0.01)
train_dataset, eval_dataset = splits["train"], splits["test"]

# Configure loss with cross-GPU gathering

train_loss = losses.Contrastive(
    model=model,
    gather_across_devices=True  # Enables global batch construction

)

```

The `gather_across_devices=True` parameter triggers the `all_gather` function in [`pylate/utils/distributed.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/distributed.py), which collects embeddings from every GPU while preserving gradients for the local rank.

### Training Arguments Configuration

```python
args = SentenceTransformerTrainingArguments(
    output_dir="output/contrastive-multi-gpu",
    num_train_epochs=1,
    per_device_train_batch_size=32,
    fp16=True,  # Mixed precision training

    run_name="contrastive-multi-gpu"
)

trainer = SentenceTransformerTrainer(
    model=model,
    args=args,
    train_dataset=train_dataset,
    eval_dataset=eval_dataset,
    loss=train_loss,
    evaluator=evaluation.ColBERTTripletEvaluator(
        anchors=eval_dataset["query"],
        positives=eval_dataset["positive"],
        negatives=eval_dataset["negative"]
    ),
    data_collator=utils.ColBERTCollator(model.tokenize)
)

trainer.train()

```

The `SentenceTransformerTrainer` automatically detects distributed environment variables and initializes Distributed Data Parallel (DDP) when launched appropriately.

## Launching Multi-GPU Training

### Method 1: Using torchrun (Recommended)

Set the required environment variables and launch:

```bash
export MASTER_ADDR=127.0.0.1
export MASTER_PORT=29500
export WORLD_SIZE=4

torchrun --nproc_per_node=4 examples/train/contrastive.py

```

`torchrun` automatically assigns `RANK` to each process. The helper in [`pylate/utils/distributed.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/distributed.py) validates the world size and pins each process to its corresponding GPU using `rank % torch.cuda.device_count()`.

### Method 2: Using Hugging Face Accelerate

Configure Accelerate for your hardware:

```bash
accelerate config

# Select "Distributed Data Parallel" when prompted

```

Launch the training script:

```bash
accelerate launch examples/train/contrastive.py

```

Accelerate generates a temporary configuration file that populates the same environment variables used by `torchrun`, ensuring seamless integration with PyLate's distributed utilities.

## Understanding the Distributed Architecture

PyLate's multi-GPU capability relies on several interconnected components:

| Component | Multi-GPU Role | Source Location |
|-----------|---------------|-----------------|
| **Process Group Initialization** | Sets up NCCL backend and assigns GPU devices per rank | [`pylate/utils/distributed.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/distributed.py) |
| **Cross-GPU Gathering** | Collects embeddings from all ranks using `all_gather` while preserving local gradients | [`pylate/utils/distributed.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/distributed.py) |
| **Contrastive Losses** | `Contrastive` and `CachedContrastive` classes that invoke `all_gather` when `gather_across_devices=True` | [`pylate/losses/contrastive.py`](https://github.com/lightonai/pylate/blob/main/pylate/losses/contrastive.py) |
| **Training Orchestration** | `SentenceTransformerTrainer` automatically handles DDP gradient synchronization | Integrated via Sentence-Transformers |

The `all_gather` implementation in [`pylate/utils/distributed.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/distributed.py) is particularly critical—it enables the contrastive loss to compute similarities against a global batch rather than just the local batch, significantly improving training stability and model quality.

## Optimization Tips and Troubleshooting

### Memory Management

If you encounter out-of-memory errors on individual GPUs:

- Increase `per_device_train_batch_size` while keeping `gather_across_devices=True`. The effective batch size becomes `batch_size × world_size`, allowing smaller per-device batches while maintaining global batch size.
- Enable gradient checkpointing by adding `model.gradient_checkpointing_enable()` before training.

### Precision and Device Compatibility

- Verify GPU compatibility for mixed precision: check `torch.cuda.get_device_capability()` returns a value ≥ 7.0 for FP16 support.
- If using older GPUs, set `fp16=False` and optionally enable `bf16=True` if supported.

### Debugging Distributed Issues

- Test your script on a single GPU first by running without `torchrun`. The `gather_across_devices` parameter will emit a warning and return local tensors only, allowing you to verify data loading and model logic.
- For custom data pipelines, explicitly call `pylate.utils.distributed.init(rank)` before data loading to ensure correct device assignment.
- When running on CPU-only clusters, modify the backend to `gloo` in the distributed initialization.

## Summary

- **Install PyLate** with development dependencies to access distributed utilities in [`pylate/utils/distributed.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/distributed.py).
- **Enable cross-GPU gathering** by setting `gather_across_devices=True` in `losses.Contrastive` or `losses.CachedContrastive`.
- **Launch training** using `torchrun --nproc_per_node=N` or `accelerate launch` to automatically initialize Distributed Data Parallel.
- **Leverage global batches** through the `all_gather` implementation, which collects embeddings from all ranks while preserving local gradients.
- **Optimize memory** by adjusting per-device batch sizes and utilizing mixed precision (`fp16=True`) to train larger effective batches across multiple GPUs.

## Frequently Asked Questions

### How does PyLate handle gradient synchronization across multiple GPUs?

PyLate relies on the Hugging Face `SentenceTransformerTrainer`, which internally uses PyTorch's Distributed Data Parallel (DDP). DDP automatically synchronizes gradients across all processes during the backward pass. Additionally, when `gather_across_devices=True` is set in the `Contrastive` loss, PyLate uses the `all_gather` function from [`pylate/utils/distributed.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/distributed.py) to collect embeddings from all GPUs before computing the loss, ensuring the model trains on a global batch rather than individual local batches.

### What is the difference between `Contrastive` and `CachedContrastive` when using multi-GPU training?

Both `Contrastive` and `CachedContrastive` classes in [`pylate/losses/contrastive.py`](https://github.com/lightonai/pylate/blob/main/pylate/losses/contrastive.py) support the `gather_across_devices` parameter for multi-GPU training. The standard `Contrastive` loss computes similarities between queries and documents in a single forward pass. `CachedContrastive` is designed for scenarios where document embeddings can be pre-computed and cached, reducing redundant computation during training. When `gather_across_devices=True`, both losses invoke the same `all_gather` utility to synchronize embeddings across the process group before computing the contrastive objective.

### Can I use PyLate multi-GPU training on a single machine with multiple GPUs?

Yes, PyLate multi-GPU training is fully supported on single-machine, multi-GPU setups using `torchrun` or `accelerate launch`. Set the `MASTER_ADDR` to `127.0.0.1` and choose an available `MASTER_PORT` (e.g., 29500). Then launch with `torchrun --nproc_per_node=4 examples/train/contrastive.py` (adjusting the number of processes to match your GPU count). The [`pylate/utils/distributed.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/distributed.py) helper will automatically detect the local rank and bind each process to its corresponding GPU using `rank % torch.cuda.device_count()`.

### How do I debug a PyLate training script before running it on multiple GPUs?

To debug a PyLate training script, run it directly with Python on a single GPU or CPU without using `torchrun` or `accelerate`. When `gather_across_devices=True` is set but no distributed environment is detected, the `all_gather` function in [`pylate/utils/distributed.py`](https://github.com/lightonai/pylate/blob/main/pylate/utils/distributed.py) will emit a warning and return the local tensor only, allowing the script to execute without errors. This enables you to verify data loading, model initialization, and training logic before scaling to multiple GPUs. Once debugging is complete, simply prepend `torchrun --nproc_per_node=N` to your command to launch distributed training.