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

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:

pip install "pylate[dev]" torch torchvision

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

pip install accelerate

PyLate’s core distributed utilities are located in 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 file demonstrates the complete setup.

Essential Code Components


# 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, which collects embeddings from every GPU while preserving gradients for the local rank.

Training Arguments Configuration

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

Set the required environment variables and launch:

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 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:

accelerate config

# Select "Distributed Data Parallel" when prompted

Launch the training script:

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
Cross-GPU Gathering Collects embeddings from all ranks using all_gather while preserving local gradients pylate/utils/distributed.py
Contrastive Losses Contrastive and CachedContrastive classes that invoke all_gather when gather_across_devices=True 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 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.
  • 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 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 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 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 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.

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 →