How PyLate Implements Late Interaction in ColBERT Models: A Deep Dive into the Source Code
PyLate implements late interaction by encoding queries and documents into token-level embeddings first, then computing similarity using the MaxSim operator—summing the maximum dot-product between each query token and all document tokens—rather than mixing query-document signals inside the transformer layers.
The lightonai/pylate repository provides a modular implementation of the ColBERT architecture, distinguishing itself from standard dense retrievers through its delayed similarity computation. This approach, known as late interaction, keeps query and document representations separate during the expensive transformer encoding phase, deferring comparison until the final scoring stage.
Token-Level Encoding with Query and Document Prefixes
At the heart of PyLate’s implementation is the ColBERT class in pylate/models/colbert.py. Unlike standard sentence embeddings that pool tokens into a single vector, PyLate preserves token-level embeddings throughout the encoding process.
The tokenize() method prepares sequences by injecting special prefix tokens that signal whether the input is a query or document. For queries, it prepends [Q]; for documents, it prepends [D]. When query expansion is enabled, the method also pads queries to a fixed length to ensure sufficient token coverage during the late interaction phase.
After tokenization, the model passes sequences through a BERT-style encoder (self._first_module()) followed by a linear projection layer (Dense). This yields token_embeddings with shape (batch, tokens, dim), where each token retains its own high-dimensional representation rather than being aggregated early.
The Late Interaction Mechanism: MaxSim Scoring
The defining characteristic of ColBERT’s architecture is the separation of encoding and interaction. In pylate/scores/similarity_functions.py, the similarity function enum maps the string "MaxSim" to the concrete implementation found in pylate/scores/scores.py.
The actual late interaction occurs in the colbert_scores function:
def colbert_scores(query_embeddings, doc_embeddings):
# query_embeddings: (Q, dim) - one embedding per query token
# doc_embeddings: (D, dim) - one embedding per document token
# Compute the dot-product matrix (Q × D)
sim = torch.mm(query_embeddings, doc_embeddings.t())
# Take the maximum similarity for each query token, then sum
max_sim, _ = sim.max(dim=1)
return max_sim.sum()
This late interaction approach builds the full dot-product matrix after encoding completes. Each query token finds its best matching document token via max(dim=1), and these maximum similarities are summed to produce the final relevance score. The similarity_fn_name property of the ColBERT class defaults to "MaxSim", ensuring this behavior is used when model.similarity() is invoked.
Masking Strategies for Clean Embeddings
PyLate applies selective masking to remove noise from the token sequences before similarity computation.
Document Skip-List Masking
For documents, the implementation uses self.skiplist_mask (defined in pylate/utils/tensor.py) to filter out punctuation and other low-information tokens. This prevents meaningless matches between query terms and document punctuation during the MaxSim operation.
Query Attention Masking
Queries receive different treatment based on the expansion setting. When query expansion is active, masks default to all-ones to preserve the padded structure. Otherwise, standard attention masks apply to ignore padding tokens naturally.
Pooling for Long Documents
To handle documents exceeding token limits without sacrificing granularity, pylate/models/colbert.py implements pool_embeddings_hierarchical. When pool_factor > 1 and the input is a document, PyLate applies Ward clustering hierarchically to the token embeddings, averages each cluster, and reduces the sequence length before the MaxSim step.
This hierarchical pooling maintains semantic coverage while drastically reducing the computational cost of the late interaction matrix multiplication for long texts.
Putting It All Together: From Encoding to Similarity
The complete late-interaction pipeline flows through the encode() method, which returns token embeddings, and the similarity() method, which delegates to the stored self._similarity function pointing to MaxSim. The retrieval wrapper in pylate/retrieve/colbert.py orchestrates this by encoding queries and documents separately, then invoking the late-interaction scorer for ranking.
from pylate import models
# Load a ColBERT model
model = models.ColBERT(
model_name_or_path="sentence-transformers/all-MiniLM-L6-v2",
device="cpu",
)
# Encode query and document (token-level embeddings)
query_emb = model.encode(
"what is the capital of france?",
is_query=True, # adds [Q] prefix and expands if needed
convert_to_numpy=False, # keep as torch tensors for scoring
)
doc_emb = model.encode(
"Paris is the capital city of France.",
is_query=False, # adds [D] prefix, no expansion
convert_to_numpy=False,
)
# Compute late-interaction score via MaxSim
score = model.similarity(query_emb, doc_emb)
print(f"Late-interaction score: {score:.4f}")
Summary
- Late interaction delays similarity computation until after independent encoding of queries and documents, preserving fine-grained token-level signals.
- PyLate implements this through the MaxSim operator in
pylate/scores/scores.py, which computes maximum dot-products between query and document tokens. - The
tokenize()method incolbert.pyhandles special prefix tokens ([Q]and[D]) and optional query expansion to ensure robust matching. - Skip-list masking removes punctuation from document embeddings, while hierarchical pooling compresses long documents via Ward clustering before scoring.
- The separation of encoding (BERT + Dense projection) and interaction (MaxSim) distinguishes PyLate from early-interaction dense retrieval models.
Frequently Asked Questions
What is late interaction in ColBERT?
Late interaction is an architectural pattern where query and document embeddings are computed independently through transformer layers, and their similarity is calculated only at the final stage using token-level comparisons. This contrasts with early-interaction models that concatenate queries and documents before the encoding phase, allowing ColBERT to scale efficiently while maintaining granular matching capabilities.
How does PyLate differ from standard BERT-based retrievers?
Standard BERT retrievers typically pool token embeddings into a single dense vector per sequence (often via [CLS] token or mean pooling) and compute similarity using cosine similarity or dot product between these single vectors. PyLate retains all token embeddings and uses the MaxSim operator, enabling finer-grained matching where specific query terms align with specific document terms rather than holistic sequence representations.
What is the MaxSim operation and why is it used?
MaxSim (Maximum Similarity) is the core scoring function that computes, for each query token, the maximum dot-product similarity against all document tokens, then sums these maxima. It is used because it captures the best possible alignment for each query term while remaining computationally tractable, effectively modeling soft term matching without requiring exact lexical overlap.
When should I use document pooling in PyLate?
Enable document pooling by setting pool_factor > 1 when indexing very long documents that would otherwise exceed token length limits or create prohibitively large similarity matrices during retrieval. The hierarchical Ward clustering in pool_embeddings_hierarchical reduces token count while preserving semantic clusters, making it ideal for retrieving from lengthy passages or entire documents rather than short snippets.
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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →