FlexAttention for Inference

PyTorch 2.7 extends FlexAttention to inference mode. This API lets you define attention variants (causal, sliding window, custom mask) via a simple Python function, compiled into an optimized CUDA kernel by torch.compile.

Practical Usage

python
import torch
from torch.nn.attention.flex_attention import (
    flex_attention, create_block_mask
)


# Define a sliding window attention mask
def sliding_window(b, h, q_idx, kv_idx):
    return (q_idx - kv_idx).abs() <= 512


# Create the optimized mask
block_mask = create_block_mask(
    sliding_window, B=1, H=1,
    Q_LEN=4096, KV_LEN=4096,
)

# Use with flex_attention
q = torch.randn(1, 8, 4096, 64, device='cuda')
k = torch.randn(1, 8, 4096, 64, device='cuda')
v = torch.randn(1, 8, 4096, 64, device='cuda')

out = flex_attention(q, k, v, block_mask=block_mask)
print(out.shape)  # torch.Size([1, 8, 4096, 64])

Sources