[Bug] Context parallel crashes with asymmetric K/V head dims (GQA + enable_mla path)
维护者通常 2 天内回复
还没有人认领这个 Issue。
- #2901 来自 @beccohov —— 已关闭,未合并
评估
- 难度
- 3/5
- 预计耗时
- 1-2 天
- 新手友好度
- 68/100
- Issue 类型
- 缺陷
- 描述清晰度
- 描述清楚
- 活跃度
- 冷清
调研方向
从 transformer_engine/pytorch/attention/dot_product_attention/context_parallel.py 开始,重点查看第 1316 行附近的 enable_mla 逻辑,以及第 1854 行附近的输出 reshape。使用不对称的 K/V 维度和 GQA 运行提供的分布式 CUDA 复现,然后比较 CP 路径与非 CP 路径。当 CP forward pass 成功完成并生成具有预期 attention-head 形状的输出时,即表示完成。
由索引模型根据 Issue 内容生成。
描述
Describe the bug
When using DotProductAttention with context parallelism (CP) and asymmetric K/V head dimensions (kv_channels=(k_dim, v_dim) where k_dim != v_dim), the CP forward pass crashes with a shape mismatch in context_parallel.py.
The root cause: enable_mla = k.shape[-1] != v.shape[-1] (here) triggers the MLA code path, which reshapes the attention output using v_shape (derived from the V input tensor with num_kv_heads). However, with GQA the attention output has num_attention_heads (post-expansion), not
num_kv_heads, causing a size mismatch.
This affects models like MiMo-V2-Flash from Megatron Bridge (I'm currently implementing support for it here) which uses head_dim=192 for Q/K but v_head_dim=128 for V with GQA (num_attention_heads=64, num_kv_heads=4).
Steps/Code to reproduce bug
import os
import torch
import transformer_engine as te
from transformer_engine.pytorch.attention.dot_product_attention import DotProductAttention
print(f"TE version: {te.__version__}")
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
torch.cuda.set_device(local_rank)
torch.distributed.init_process_group("nccl")
device = torch.device(f"cuda:{local_rank}")
B, T, num_heads, num_kv_heads = 1, 32, 8, 2
qk_head_dim, v_head_dim = 64, 48 # asymmetric: k != v
cp_group = torch.distributed.group.WORLD
# CP case doesn't work:
attn = DotProductAttention(
num_attention_heads=num_heads,
kv_channels=(qk_head_dim, v_head_dim),
num_gqa_groups=num_kv_heads,
attention_dropout=0.0,
tp_size=1,
tp_group=None,
cp_global_ranks=list(range(torch.distributed.get_world_size())),
cp_group=cp_group,
cp_stream=torch.cuda.Stream(),
softmax_type="vanilla", #"learnable" SWA uses attention sink bias
).to(device)
# non-CP works:
# attn = DotProductAttention(
# num_attention_heads=num_heads,
# kv_channels=(qk_head_dim, v_head_dim),
# num_gqa_groups=num_kv_heads,
# attention_dropout=0.0,
# softmax_type="vanilla", #"learnable" SWA uses attention sink bias
# ).to(device)
q = torch.randn(T, B, num_heads, qk_head_dim, device=device, dtype=torch.bfloat16)
k = torch.randn(T, B, num_kv_heads, qk_head_dim, device=device, dtype=torch.bfloat16)
v = torch.randn(T, B, num_kv_heads, v_head_dim, device=device, dtype=torch.bfloat16)
out = attn(q, k, v, attn_mask_type="causal")
Expected behavior
CP forward should handle asymmetric K/V head dims with GQA correctly. The output reshape at context_parallel.py should use the attention output's actual shape (num_attention_heads) rather than the V input shape (num_kv_heads).
Environment overview (please complete the following information)
- PyTorch: 2.7
- Transformer Engine: 2.13.0
- CUDA: 13.0
Device details
- GPU model
Additional context
Add any other context about the problem here.
- 主要语言
- Python
- 星标
- 3.6k
- 派生
- 851
- 平均合并
- 5 天 1 小时
- 30 天内合并 PR
- 52
环境准备
- 没有 Dockerfile 或 Docker Compose 文件
- 有 Pull Request 模板
- 阅读贡献指南
从这里开始
- 先读完整个 Issue,再读项目的贡献指南。
- 在 Issue 下留言说明你要接手 —— 这能避免两个人做同样的事。
- Fork 仓库,在一个分支上完成修改。
- 提交 Pull Request,并在描述里引用这个 Issue 编号。
NVIDIA/TransformerEngine 的其他 Issue
-
[Bug] group_quantize_fp8_blockwise: mbarrier invalidated before other threads finish waiting on it未关闭
难度 2/5 1-3 小时 新手友好度 82/100
NVIDIA/TransformerEngine#3647 ·
维护者通常 2 天内回复
-
[PyTorch] fp8_cs_quantize fake implementation returns a vector inverse scale instead of a scalar可能已有人在做 @sanjana658 今天认领。 未关闭
难度 2/5 1-3 小时 新手友好度 82/100
NVIDIA/TransformerEngine#3636 · 2 条评论 ·
维护者通常 2 天内回复
-
[Bug] Backend selection picks FA3 for training with head_dim_qk=192 / v_head_dim=128, but FA3 backward cannot run it可能已有人在做 @yuweih205 于 31 天前认领。 未关闭attention
难度 2/5 1-3 小时 新手友好度 85/100
NVIDIA/TransformerEngine#3481 · 4 条评论 ·
维护者通常 2 天内回复
-
bug
难度 2/5 1-3 小时 新手友好度 68/100
NVIDIA/TransformerEngine#2189 · 7 条评论 · 5 个 reaction ·
维护者通常 2 天内回复
-
[PyTorch] CUDA graph RNG registration floods training logs on automatic-registration builds可能已有人在做 @ksivaman 今天认领。 未关闭
难度 4/5 3-5 天 新手友好度 50/100
NVIDIA/TransformerEngine#3645 · 1 条评论 · 已指派 1 人 ·
维护者通常 2 天内回复
查看 NVIDIA/TransformerEngine 的全部 Issue
相似的 Issue
-
bug ready for review
难度 2/5 1-3 小时 新手友好度 86/100
odysseus-dev/odysseus#6641 ·
维护者通常 1 天内回复
-
bug
难度 2/5 1-3 小时 新手友好度 76/100
happypawspillaro/happypaws#78 ·
维护者通常 4 天内回复
-
pydanty:is-working
难度 2/5 1-3 小时 新手友好度 82/100
pydantic/pydantic-ai#10020 ·
维护者通常 1 天内回复
-
Bug
难度 2/5 1-3 小时 新手友好度 78/100
ansible-collections/ibm_zos_core#2650 ·
-
hw: pvc tests: vllm vllm
难度 2/5 1-3 小时 新手友好度 68/100
intel/intel-xpu-backend-for-triton#8362 ·
维护者通常 1 天内回复