PyLate Attention Implementations: Eager, SDPA, and Flash Attention 2 Explained

PyLate supports three attention implementations—"eager", "sdpa", and "flash_attention_2"—allowing you to optimize ColBERT model performance based on your PyTorch version and GPU hardware.

The ColBERT class in the lightonai/pylate repository inherits from SentenceTransformer and exposes an attn_implementation parameter. This parameter controls which attention kernel is used during inference, directly impacting memory efficiency and computational speed. You can specify the implementation explicitly or let PyLate automatically select the best available option for your environment.

Available Attention Implementations in PyLate

PyLate delegates attention computation to one of three backends, each with distinct performance characteristics and hardware requirements.

Eager (Manual PyTorch Implementation)

The "eager" option uses a manual, pure-PyTorch implementation of scaled dot-product attention. This is the most compatible option, requiring no special libraries or specific GPU architectures.

  • Best for: CPU inference, debugging, or environments where specialized kernels are unavailable.
  • Performance: Baseline speed; generally slower than optimized kernels on GPU.

SDPA (Scaled Dot-Product Attention)

The "sdpa" implementation leverages torch.nn.functional.scaled_dot_product_attention, available in PyTorch 2.1 and later. This kernel automatically selects the most efficient algorithm (FlashAttention, memory-efficient attention, or the math implementation) based on your hardware and input dimensions.

  • Best for: General GPU inference with PyTorch ≥2.1.
  • Performance: Significantly faster than eager on compatible GPUs, with optimized memory access patterns.

Flash Attention 2

The "flash_attention_2" option uses the high-performance FlashAttention 2 library from Dao-AILab. This implementation requires the flash-attn package to be installed and a compatible NVIDIA GPU (Ampere, Ada Lovelace, or Hopper architecture).

  • Best for: Maximum throughput on modern NVIDIA GPUs with long sequences.
  • Performance: Lowest memory footprint and highest speed for large batch sizes and long contexts.

Automatic Selection Behavior

If you do not specify attn_implementation, PyLate automatically selects the most efficient available backend. According to the implementation in pylate/models/colbert.py (lines 119-124), the selection logic prefers SDPA when PyTorch 2.1.1 or later is detected. If SDPA is unavailable, the system falls back to the "eager" implementation.

This automatic selection ensures optimal performance without requiring manual configuration, while still allowing advanced users to override the choice for specific hardware or debugging needs.

Configuring Attention Implementations in Code

You specify the attention backend via the attn_implementation parameter when instantiating the ColBERT class. This value is passed through model_kwargs to the underlying Hugging Face transformer model.

Using Eager Mode

from pylate import models

# Force manual PyTorch attention (good for CPU or debugging)

model = models.ColBERT(
    model_name_or_path="sentence-transformers/all-MiniLM-L6-v2",
    attn_implementation="eager",
    device="cpu",
)

Using SDPA

from pylate import models

# Use PyTorch's native scaled_dot_product_attention (requires torch>=2.1)

model = models.ColBERT(
    model_name_or_path="sentence-transformers/all-MiniLM-L6-v2",
    attn_implementation="sdpa",
    device="cuda",
)

Using Flash Attention 2

from pylate import models

# Use FlashAttention 2 for maximum GPU efficiency

model = models.ColBERT(
    model_name_or_path="sentence-transformers/all-MiniLM-L6-v2",
    attn_implementation="flash_attention_2",
    device="cuda",
)

Implementation Details and Source Code

The attention implementation configuration is handled in pylate/models/colbert.py. The ColBERT class accepts attn_implementation as an initialization argument and passes it to the underlying transformer model through the model_kwargs dictionary.

As documented in the source (lines 119-124), the supported string values are:

  • "eager" – manual implementation of attention
  • "sdpa" – calls torch.nn.functional.scaled_dot_product_attention
  • "flash_attention_2" – wraps the Dao-AILab FlashAttention 2 library

The pylate/utils/collator.py module handles attention mask creation for query and document batches, ensuring that the selected attention kernel receives properly formatted inputs during contrastive training. The pylate/losses/contrastive.py module then utilizes these attention masks when computing the contrastive loss across the different implementation backends.

Summary

  • PyLate supports three attention implementations: "eager" (manual PyTorch), "sdpa" (PyTorch 2.1+ optimized), and "flash_attention_2" (Dao-AILab library).
  • The attn_implementation parameter in ColBERT controls which backend is used, with values passed through to the underlying Hugging Face transformer.
  • Automatic selection defaults to SDPA when available (PyTorch ≥2.1.1), falling back to eager otherwise.
  • Flash Attention 2 requires specific GPU architectures (Ampere or newer) and the flash-attn package installation.

Frequently Asked Questions

What is the default attention implementation in PyLate?

If you do not specify attn_implementation, PyLate automatically selects the best available option. According to the source code in pylate/models/colbert.py, the system prefers "sdpa" when PyTorch 2.1.1 or later is installed. If SDPA is unavailable, it falls back to the "eager" implementation.

Does PyLate support Flash Attention on all GPUs?

No, Flash Attention 2 requires specific hardware. You need an NVIDIA GPU with compute capability 8.0 or higher (Ampere, Ada Lovelace, or Hopper architectures). Additionally, you must install the flash-attn library separately, as it is not included in the base PyLate dependencies.

How do I check which attention implementation is currently active?

When you instantiate a ColBERT model without specifying attn_implementation, PyLate selects the backend automatically based on your PyTorch version. To verify which implementation is loaded, you can inspect the model configuration or check the PyTorch version (torch.__version__)—if it is 2.1.1 or higher and you did not override the setting, SDPA is being used.

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 →