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:

  1. pool_factor > 1: The reduction factor determines the number of clusters (num_clusters = max(num_embeddings // pool_factor, 1))
  2. 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:

  1. Protection: The first protected_tokens (default 1) embeddings are preserved unchanged, typically safeguarding the [CLS] token or special markers
  2. Distance Calculation: Remaining tokens are converted to a cosine-distance matrix
  3. Clustering: Ward’s hierarchical agglomerative clustering groups tokens into max(num_embeddings // pool_factor, 1) clusters
  4. Averaging: Mean vectors are computed for each cluster, creating representative embeddings
  5. 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 > 1 and is_query=False in encode() or encode_multi_process().
  • The pool_embeddings_hierarchical method in pylate/models/colbert.py handles the clustering logic, preserving protected_tokens before 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:

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 →