How to Utilize Hierarchical Pooling with pool_embeddings_hierarchical in PyLate
Use ColBERT.pool_embeddings_hierarchical() in PyLate to reduce long document token embeddings into compact clusters by setting pool_factor > 1 and is_query=False during encoding.
PyLate provides a specialized hierarchical pooling mechanism for its ColBERT implementation that compresses lengthy document representations while preserving semantic fidelity. This technique clusters token-level embeddings using Ward’s method and averages vectors within each cluster, significantly reducing storage and computation costs for retrieval pipelines.
Understanding Hierarchical Pooling in PyLate
The pool_embeddings_hierarchical Method
The core functionality resides in pylate/models/colbert.py (lines 994-1014). The method signature is:
ColBERT.pool_embeddings_hierarchical(
self,
documents_embeddings: list[torch.Tensor],
pool_factor: int = 1,
protected_tokens: int = 1,
) → list[torch.Tensor]
This method accepts a list of token-level embedding tensors (one per document) and returns a list of pooled tensors with reduced dimensionality along the sequence length axis.
When Hierarchical Pooling Activates
According to the source code in pylate/models/colbert.py (lines 743-749), hierarchical pooling automatically triggers during the encoding pipeline when two conditions are met:
pool_factor > 1: The reduction factor determines the number of clusters (num_clusters = max(num_embeddings // pool_factor, 1))is_query=False: The pooling applies only to documents; queries retain their full token-level representations for fine-grained late interaction
How to Use pool_embeddings_hierarchical
Direct Method Invocation
For advanced use cases where you already possess token-level embeddings, invoke the method directly:
from pylate.models.colbert import ColBERT
import torch
# Initialize model
model = ColBERT("bert-base-uncased")
# Example: Two documents with 120 and 95 tokens respectively
doc1_embeddings = torch.randn(120, model.hidden_dim)
doc2_embeddings = torch.randn(95, model.hidden_dim)
documents = [doc1_embeddings, doc2_embeddings]
# Apply hierarchical pooling with factor 4
pooled = model.pool_embeddings_hierarchical(
documents_embeddings=documents,
pool_factor=4, # Reduces 120 tokens to ~30 clusters
protected_tokens=1, # Preserves the first token (e.g., [CLS])
)
print([t.shape for t in pooled]) # Output: [torch.Size([31, 768]), torch.Size([24, 768])]
End-to-End Document Encoding
The typical workflow uses encode_multi_process with automatic pooling:
from pylate.models.colbert import ColBERT
model = ColBERT("bert-base-uncased")
documents = [
"The quick brown fox jumps over the lazy dog.",
"Machine learning enables computers to learn from data without explicit programming."
]
# Encode with hierarchical pooling activated
doc_embeddings = model.encode_multi_process(
sentences=documents,
pool_factor=4, # Triggers pool_embeddings_hierarchical internally
is_query=False, # Required for documents
batch_size=8,
padding=False, # Padding handled post-pooling
)
# Result: List of tensors with reduced sequence length
print(f"Document 1 pooled shape: {doc_embeddings[0].shape}")
Single-Process Encoding
For single-GPU or CPU environments, use the standard encode method with identical parameters:
doc_embeddings = model.encode(
sentences=documents,
is_query=False,
pool_factor=4, # Activates hierarchical pooling
)
Algorithm and Parameters
The hierarchical pooling algorithm in pylate/models/colbert.py follows these steps:
- Protection: The first
protected_tokens(default 1) embeddings are preserved unchanged, typically safeguarding the[CLS]token or special markers - Distance Calculation: Remaining tokens are converted to a cosine-distance matrix
- Clustering: Ward’s hierarchical agglomerative clustering groups tokens into
max(num_embeddings // pool_factor, 1)clusters - Averaging: Mean vectors are computed for each cluster, creating representative embeddings
- Reconstruction: Protected tokens are prepended to the pooled clusters, returning the final tensor
| Parameter | Effect |
|---|---|
pool_factor |
Determines compression ratio. Higher values create fewer clusters (e.g., pool_factor=4 reduces 100 tokens to ~25). |
protected_tokens |
Number of initial tokens excluded from clustering. Default of 1 preserves the first token for semantic anchoring. |
Summary
- Hierarchical pooling in PyLate compresses document token embeddings via Ward’s clustering, reducing storage and computational overhead for retrieval.
- Activate pooling by setting
pool_factor > 1andis_query=Falseinencode()orencode_multi_process(). - The
pool_embeddings_hierarchicalmethod inpylate/models/colbert.pyhandles the clustering logic, preservingprotected_tokensbefore averaging embeddings within each cluster. - Use direct method invocation only when working with pre-computed token embeddings; otherwise, rely on the high-level encoding API.
Frequently Asked Questions
What is the optimal pool_factor for document retrieval?
The optimal pool_factor depends on your latency and accuracy requirements. A value of 4 provides a 4x reduction in sequence length with minimal impact on retrieval accuracy for most datasets, while values above 8 may degrade fine-grained matching performance. Benchmark on your specific corpus to determine the best trade-off.
Can I use hierarchical pooling for queries instead of documents?
No. According to the implementation in pylate/models/colbert.py, hierarchical pooling is disabled for queries (is_query=True). ColBERT’s late interaction mechanism requires full token-level query representations to compute fine-grained similarity with pooled document embeddings. Queries should remain un-pooled to maintain retrieval accuracy.
Why are protected_tokens necessary?
The protected_tokens parameter preserves the first N tokens (typically the [CLS] token) from clustering, ensuring that critical semantic markers or special tokens retain their original embeddings. This is crucial because the first token often carries document-level classification signals that Ward’s clustering might otherwise dilute through averaging with unrelated token clusters.
How does hierarchical pooling affect inference performance?
Hierarchical pooling significantly improves inference performance for long documents by reducing the number of vectors stored in the index and compared during similarity search. With pool_factor=4, you achieve approximately 4x reduction in memory usage and faster MaxSim operations, though the clustering step adds minor overhead during encoding (typically <5% of total encoding time).
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 →