[Bug] Backend selection picks FA3 for training with head_dim_qk=192 / v_head_dim=128, but FA3 backward cannot run it
Los mantenedores suelen responder en 2 días
Evaluación
- Dificultad
- 2/5
- Tiempo estimado
- 1-3 horas
- Aptitud para principiantes
- 85/100
- Tipo de issue
- Error
- Claridad
- Bien especificado
- Estado de actividad
- Activo
- Stack tecnológico
- python
- Área
- machine-learning
Línea de trabajo
Comienza en dot_product_attention/utils.py, en _is_fa3_supported(), y compara sus comprobaciones de capacidades con el caso de entrenamiento descrito en el issue. Ejecuta repro.py con la configuración de entorno indicada para confirmar la selección del backend y el fallo en el backward. El trabajo estará terminado cuando el entrenamiento evite la geometría FA3 no compatible y seleccione un backend viable o indique que no hay ninguno disponible.
Escrito por el modelo de indexación a partir del texto del issue.
Descripción
Describe the bug
With an MLA-style geometry (head_dim_qk=192, v_head_dim=128, GQA), when FusedAttention is unavailable (e.g. NVTE_FUSED_ATTN=0, which Megatron-LM sets when launched with --attention-backend flash), TE selects FlashAttention 3 in training mode. The forward pass succeeds; the backward crashes deep in the kernel:
[DEBUG | DotProductAttention]: Disabling FusedAttention due to NVTE_FUSED_ATTN=0
[DEBUG | DotProductAttention]: Disabling FlashAttention 2 as it does not support MLA.
[DEBUG | DotProductAttention]: Selected backend = FlashAttention (3.0.0b1)
...forward OK...
File "flash_attn_3/flash_attn_interface.py", line 123, in _flash_attn_backward
RuntimeError: out must have shape (batch_size, seqlen_q, num_heads, head_size)
Both plain causal and sliding-window (window_size=(127, 0)) hit the same crash.
Root cause (as far as I can tell)
_is_fa3_supported() in dot_product_attention/utils.py allows head_dim_qk != head_dim_v when 128 < qk <= 192 and 96 < v <= 128 — this matches FA3's forward support, but the function never consults is_training. FA3's backward for hdimQK=192/hdimV=128 is not implemented (open feature request: Dao-AILab/flash-attention#1487). So the unsupported-backward config passes selection and only fails mid-backward with an opaque shape error.
To Reproduce
# NVTE_FLASH_ATTN=1 NVTE_FUSED_ATTN=0 NVTE_DEBUG=1 NVTE_DEBUG_LEVEL=2 python repro.py
import torch
from transformer_engine.pytorch import DotProductAttention
HQ, KV, DQK, DV, S = 16, 1, 192, 128, 4096
q = torch.randn(S, 1, HQ, DQK, dtype=torch.bfloat16, device="cuda", requires_grad=True)
k = torch.randn(S, 1, KV, DQK, dtype=torch.bfloat16, device="cuda", requires_grad=True)
v = torch.randn(S, 1, KV, DV, dtype=torch.bfloat16, device="cuda", requires_grad=True)
dpa = DotProductAttention(num_attention_heads=HQ, kv_channels=(DQK, DV), num_gqa_groups=KV,
attention_dropout=0.0, qkv_format="sbhd", attn_mask_type="causal").cuda().train()
out = dpa(q, k, v) # forward OK
out.sum().backward() # RuntimeError: out must have shape (batch_size, seqlen_q, num_heads, head_size)
Expected behavior
During training, FA3 should be filtered out for geometries whose backward it cannot run — same as other capability filters — so selection either falls back to a viable backend or fails fast with a clear "no viable backend" error, instead of crashing inside flash_attn_3_cuda.bwd after a successful forward.
Environment
TE 2.10.0+769ed778 · torch 2.9.1+cu130 · flash-attn 2.7.4.post1 · flash_attn_3 3.0.0b1 · H200 (sm90) · CUDA 13.0
- Lenguaje dominante
- Python
- Estrellas
- 3.6k
- Forks
- 844
- Merge medio
- 4 d 42 min
- PR fusionados (30 d)
- 49
Preparar el entorno
- Sin Dockerfile ni archivo de Docker Compose
- Tiene una plantilla de pull request
- Leer la guía de contribución
Primeros pasos
- Lee el issue completo y luego la guía de contribución del proyecto.
- Comenta en el issue que vas a ocuparte — evita que dos personas hagan lo mismo.
- Haz un fork del repositorio y trabaja en una rama.
- Abre un pull request que haga referencia al número del issue.
Más de NVIDIA/TransformerEngine
-
Increase MAX_TENSOR_NUMAbiertobug
Dificultad 2/5 1-3 horas Aptitud para principiantes 68/100
NVIDIA/TransformerEngine#2189 · 7 comentarios · 5 reacciones ·
Los mantenedores suelen responder en 2 días
-
[BUG] Grouped MXFP8 quantization is not concurrency safe with multiple streamsPosiblemente ocupada @kainzhong la tomó hoy. Abiertobug
Dificultad 4/5 3-5 días Aptitud para principiantes 25/100
NVIDIA/TransformerEngine#3630 ·
Los mantenedores suelen responder en 2 días
-
[bug] NVFP4 + `torch.compile`: errors with 3D inputPosiblemente ocupada @pggPL la tomó hace 1 día. Abiertobug
Dificultad 2/5 1-3 horas Aptitud para principiantes 25/100
NVIDIA/TransformerEngine#3626 ·
Los mantenedores suelen responder en 2 días
-
Multi-tensor swizzle kernels fail with "too many resources requested for launch" (missing __launch_bounds__)Posiblemente ocupada @ravimajeti la tomó hace 2 días. Abierto
Dificultad 2/5 1-3 horas Aptitud para principiantes 25/100
NVIDIA/TransformerEngine#3621 ·
Los mantenedores suelen responder en 2 días
-
[PyTorch] Avoid selecting FA4 for deterministic training on SM120Posiblemente ocupada Un pull request vinculado a esta issue está abierto o ya se fusionó. Abiertobug
Dificultad 3/5 1-2 días Aptitud para principiantes 35/100
NVIDIA/TransformerEngine#3594 ·
Los mantenedores suelen responder en 2 días
Todos los issues de NVIDIA/TransformerEngine
Issues similares
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 72/100
Juniper/ansible-junos-stdlib#904 ·
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 82/100
pollen-robotics/reachy_mini#1457 ·
Los mantenedores suelen responder en 1 día
-
area:runtime good first issue
Dificultad 2/5 1-3 horas Aptitud para principiantes 72/100
WATonomous/wato_f1tenth#39 ·
-
Dificultad 2/5 1-3 horas Aptitud para principiantes 82/100
FireDynamics/fdsreader#123 ·
-
Dificultad 1/5 Menos de una hora Aptitud para principiantes 85/100