Techniques for Handling Long Documents Effectively in PyLate

PyLate prevents memory overflow when processing extensive corpora through chunked encoding, batched similarity computation, and device-aware multi-process pools, enabling scalable neural retrieval across millions of tokens.

PyLate is a neural retrieval library designed to scale ColBERT-style models to massive text corpora. When documents grow lengthy or batch sizes increase, naive processing quickly exhausts GPU memory and system RAM. The library implements sophisticated chunking strategies for handling long documents effectively in PyLate, distributing computational load both horizontally across devices and vertically within operations.

Chunked Sentence Encoding

Automatic Chunk Sizing in encode_multi_process

The encode_multi_process method in pylate/models/colbert.py (lines 1010-1025) implements intelligent workload distribution to keep per-process memory bounded. Rather than loading entire corpora into worker memory, the system calculates an optimal chunk size based on the total sentence count and available worker processes.


# pylate/models/colbert.py

if chunk_size is None:
    chunk_size = min(
        math.ceil(len(sentences) / len(pool["processes"]) / 10), 5000
    )

logger.debug(
    f"Chunk data into {math.ceil(len(sentences) / chunk_size)} packages of size {chunk_size}"
)

The default formula allocates approximately ten chunks per available process, capped at 5000 sentences per chunk. This prevents any single worker from accumulating excessive embeddings in memory while maintaining efficient queue utilization.

Queue-Based Distribution and Ordering

After determining the chunk size, the method iterates through the input sentences, buffering them until reaching the threshold. Each buffer is packaged with its metadata and placed on the pool's input_queue. Workers consume these packages independently, process them through the ColBERT encoder, and return embeddings via the output_queue. The main thread reassembles outputs in their original order, ensuring deterministic results despite parallel execution.

Chunked Score Computation

Memory-Bounded Similarity Matrices

When training or scoring long documents, PyLate avoids materializing full similarity matrices through the CachedContrastive loss in pylate/losses/cached_contrastive.py. The _prepare_caches method (lines 78-80) implements double-chunking: it processes anchor embeddings in mini-batches while computing their similarity against all other embeddings in the collection.


# pylate/losses/cached_contrastive.py

for begin in tqdm.trange(
    0,
    batch_size,
    self.mini_batch_size,
    desc="Preparing caches",
    disable=not self.show_progress_bar,
):
    end = begin + self.mini_batch_size
    # … compute scores for a mini‑batch of anchors vs. all other embeddings …

    scores = torch.cat([...], dim=1)
    loss_mbatch = F.cross_entropy(
        input=scores / self.temperature,
        target=labels[begin:end],
        reduction="sum",
    )

The mini_batch_size parameter controls the maximum number of anchors scored against the full collection simultaneously. This vertical slicing keeps intermediate tensors linear in size rather than quadratic, preventing out-of-memory errors even with massive document sets.

Multi-Process Pool Architecture

Device-Aware Worker Distribution

PyLate's multi-processing infrastructure in pylate/utils/multi_process.py (lines 13-71 and 74-112) creates isolated worker processes that each manage their own device context. The _start_multi_process_pool function spawns processes using the 'spawn' start method for CUDA compatibility, with each worker loading the model onto its assigned GPU or CPU core.


# pylate/utils/multi_process.py

def _start_multi_process_pool(model, target_devices: list[str] = None) -> dict:
    # … pick CUDA/NPU/CPU devices …

    for device_id in target_devices:
        p = ctx.Process(
            target=_encode_multi_process_worker,
            args=(device_id, model, input_queue, output_queue),
            daemon=True,
        )
        p.start()
        processes.append(p)
    return {"input": input_queue, "output": output_queue, "processes": processes}

Each worker executes _encode_multi_process_worker, which receives a copy of the model (moved to CPU and share_memory()), loads its device-specific portion, and processes one chunk at a time. This design isolates memory per process, making it safe to run many workers on a single machine without the fragmentation typical of shared-memory threading.

Practical Implementation Examples

Encoding a Massive Corpus

To process millions of long documents, initialize a multi-process pool and leverage automatic chunking:

from pylate import models

# Initialize ColBERT model

model = models.ColBERT(
    "sentence-transformers/all-MiniLM-L6-v2",
    device="cpu",  # Workers will move to their assigned devices

)

# Start pool using all available GPUs

pool = model.start_multi_process_pool()

# Encode long documents with custom chunk size

embeddings = model.encode_multi_process(
    sentences=long_documents,  # List of lengthy text passages

    pool=pool,
    batch_size=32,           # Small per-process batch

    chunk_size=2000,         # Override default for memory constraints

    is_query=False,
)

model.stop_multi_process_pool(pool)

The chunk_size argument controls how many sentences are sent to each worker at once. Larger values reduce communication overhead but increase per-worker memory consumption. For documents exceeding typical sequence lengths, combine this with passage-level chunking before encoding.

Training with Cached Contrastive Loss

When fine-tuning on long documents, configure the CachedContrastive loss to prevent similarity matrix overflow:

from pylate import losses, models
from torch.utils.data import DataLoader

model = models.ColBERT("sentence-transformers/all-MiniLM-L6-v2")
criterion = losses.CachedContrastive(
    model=model,
    temperature=0.05,
    mini_batch_size=256,   # Controls score-matrix chunk size

)

loader = DataLoader(dataset, batch_size=64, collate_fn=utils.collator)

for batch in loader:
    loss = criterion(
        sentence_features=batch["features"],
        labels=batch["labels"],
    )
    loss.backward()
    optimizer.step()
    optimizer.zero_grad()

The mini_batch_size parameter inside CachedContrastive determines how many anchor embeddings are scored against the whole collection at once. For very long documents where embeddings contain many token-level vectors, reducing this value keeps GPU memory usage linear rather than quadratic.

Key Source Files

File Primary Role Direct Link
pylate/models/colbert.py – core model + encode_multi_process Chunked encoding logic, chunk_size handling View on GitHub
pylate/losses/cached_contrastive.py – memory‑efficient contrastive loss Double‑chunked score computation (self.mini_batch_size) View on GitHub
pylate/utils/multi_process.py – worker pool implementation Device‑aware multi‑process pool, queues View on GitHub
pylate/indexes/stanford_nlp/modeling/colbert.py – scoring primitives colbert_score_reduce masks out padding for long passages View on GitHub

Summary

PyLate tackles long-document workloads by splitting work both horizontally (across processes/devices) and vertically (inside each operation).

  • The chunked encoder in pylate/models/colbert.py keeps per-process memory low by distributing sentences in calculated chunks across worker pools.
  • Chunked score computation via CachedContrastive prevents the similarity matrix from exploding by processing anchors in mini_batch_size slices.
  • The multi-process pool isolates each chunk on a dedicated device through pylate/utils/multi_process.py, leveraging all available hardware from single GPUs to multi-node clusters.

By tuning chunk_size for encoding and mini_batch_size for training, you can adapt PyLate to handle documents of arbitrary length on any hardware configuration.

Frequently Asked Questions

What is the default chunk_size in PyLate's encode_multi_process?

The default chunk size is calculated as min(math.ceil(len(sentences) / len(pool["processes"]) / 10), 5000), as implemented in pylate/models/colbert.py. This formula allocates approximately ten chunks per available process while capping individual chunks at 5000 sentences to prevent memory accumulation in any single worker.

How does PyLate prevent out-of-memory errors during similarity scoring?

PyLate prevents OOM errors through double-chunking in the CachedContrastive loss class located in pylate/losses/cached_contrastive.py. The _prepare_caches method processes anchor embeddings in vertical slices determined by mini_batch_size, computing similarity against the full collection in mini-batches rather than materializing the entire score matrix simultaneously.

Can PyLate handle documents longer than the model's maximum sequence length?

PyLate's chunking mechanisms manage memory distribution across processes, but they do not override transformer sequence length limits. Documents exceeding the base model's maximum token limit (typically 512 or 1024 tokens) must be pre-segmented into passages before encoding. Once segmented, PyLate's chunked encoding efficiently processes these passage-level embeddings across distributed workers.

What hardware configurations work best for PyLate's multi-process encoding?

PyLate's multi-process architecture in pylate/utils/multi_process.py automatically detects and utilizes all available CUDA devices, falling back to CPU cores when GPUs are unavailable. Multi-GPU configurations achieve optimal throughput for long documents by distributing chunks across isolated device contexts. For single-GPU setups with limited VRAM, reducing chunk_size and batch_size parameters allows processing of arbitrarily long documents through sequential chunking.

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 →