[PyTorch] Integrate cuDNN GQA + DSA backend into DotProductAttention
I maintainer di solito rispondono entro 2 giorni
@cyanguwa ci sta già lavorando.
Dal 26/5/2026.
Valutazione
Questa issue non è ancora stata valutata.
Descrizione
Is your feature request related to a problem? Please describe.
Transformer Engine currently does not expose a path that combines Grouped Query Attention (GQA) with DeepSeek-style sparse attention (DSA), where each query token attends only to a TopK subset of key/value tokens. Several training workloads need this combination — a GQA attention shape (many query heads sharing fewer K/V heads) with a sparsity pattern that drops attention to all but a small index list per query. Without a TE-native backend, teams either fall back to community Triton kernels, which can't reach production-scale performance, or implement sparse attention outside of TE — losing autograd integration, kernel fusion, and parity with TE's existing attention features.
Describe the solution you'd like
Add a cuDNN-backed sparse-attention path inside DotProductAttention for the PyTorch frontend that:
- Recognizes a sparse-attention mode and dispatches to the new cuDNN GQA + DSA kernel
- Accepts a per-query sparse_indices tensor of shape [B, S_q, topk] selecting which K/V positions each query attends to
- Supports the standard GQA shape (num_attention_heads ≠ num_gqa_groups)
- Supports BF16 attention at minimum (FP8 indexer extension as a follow-on if needed)
- Integrates cleanly with TE's autograd and existing context-parallelism path
- Ships with numerical-equivalence tests against a reference dense-attention baseline restricted to the same TopK indices
cc: @cyanguwa
- Lingua principale
- Python
- Stelle
- 3.6k
- Fork
- 851
- Merge medio
- 4g 20h
- PR unite (30g)
- 56
Preparare l'ambiente
- Nessun Dockerfile né file Docker Compose
- Ha un modello di pull request
- Leggi la guida per i contributori
Come iniziare
- Leggi tutta la issue e poi la guida ai contributi del progetto.
- Commenta sulla issue per dire che te ne occupi tu — evita che due persone facciano lo stesso lavoro.
- Fai un fork del repository e lavora su un branch.
- Apri una pull request che faccia riferimento al numero della issue.
Altre issue di NVIDIA/TransformerEngine
-
[Bug] group_quantize_fp8_blockwise: mbarrier invalidated before other threads finish waiting on itAperta
Difficoltà 2/5 1-3 ore Idoneità per principianti 82/100
NVIDIA/TransformerEngine#3647 ·
I maintainer di solito rispondono entro 2 giorni
-
[PyTorch] fp8_cs_quantize fake implementation returns a vector inverse scale instead of a scalarForse già presa @sanjana658 l’ha presa 2 giorni fa. Aperta
Difficoltà 2/5 1-3 ore Idoneità per principianti 82/100
NVIDIA/TransformerEngine#3636 · 2 commenti ·
I maintainer di solito rispondono entro 2 giorni
-
[Bug] Backend selection picks FA3 for training with head_dim_qk=192 / v_head_dim=128, but FA3 backward cannot run itForse già presa @yuweih205 l’ha presa 33 giorni fa. Apertaattention
Difficoltà 2/5 1-3 ore Idoneità per principianti 85/100
NVIDIA/TransformerEngine#3481 · 4 commenti ·
I maintainer di solito rispondono entro 2 giorni
-
Increase MAX_TENSOR_NUMApertabug
Difficoltà 2/5 1-3 ore Idoneità per principianti 68/100
NVIDIA/TransformerEngine#2189 · 7 commenti · 5 reazioni ·
I maintainer di solito rispondono entro 2 giorni
-
Fused gemm + comm for CP A2A on BlackwellForse già presa @cyanguwa l’ha presa 1 giorno fa. Aperta2.22 attention
NVIDIA/TransformerEngine#3664 · 1 assegnatario ·
I maintainer di solito rispondono entro 2 giorni
Tutte le issue di NVIDIA/TransformerEngine
Issue simili
-
area:space-accuracy good first issue track:data
Difficoltà 2/5 1-3 ore Idoneità per principianti 85/100
Sara-Managed-Projects/space-radar#904 ·
I maintainer di solito rispondono entro 1 giorno