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])
