FlexAttention pour l'inférence

PyTorch 2.7 étend FlexAttention au mode inférence. Cette API permet de définir des variantes d'attention (causale, fenêtrée, avec masque personnalisé) via une simple fonction Python, compilée en un kernel CUDA optimisé par torch.compile.

Utilisation pratique

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


# Définir un masque d'attention fenêtré
def sliding_window(b, h, q_idx, kv_idx):
    return (q_idx - kv_idx).abs() <= 512


# Créer le masque optimisé
block_mask = create_block_mask(
    sliding_window, B=1, H=1,
    Q_LEN=4096, KV_LEN=4096,
)

# Utiliser avec 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