Hacktoberfest 2026:维护者为十月标记出来的 issue,仍然开放、适合新手。 浏览 Hacktoberfest issue

[Bug] Context parallel crashes with asymmetric K/V head dims (GQA + enable_mla path)

未关闭
#2,868 3 条评论 0 个 reaction 已指派 0 人 在 GitHub 查看

维护者通常 2 天内回复

还没有人认领这个 Issue。

  • #2901 来自 @beccohov —— 已关闭,未合并

评估

难度
3/5
预计耗时
1-2 天
新手友好度
68/100
Issue 类型
缺陷
描述清晰度
描述清楚
活跃度
冷清
技术栈
python, pytorch

调研方向

从 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 内容生成。

描述

bug

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

环境准备

从这里开始

  1. 先读完整个 Issue,再读项目的贡献指南。
  2. 在 Issue 下留言说明你要接手 —— 这能避免两个人做同样的事。
  3. Fork 仓库,在一个分支上完成修改。
  4. 提交 Pull Request,并在描述里引用这个 Issue 编号。

NVIDIA/TransformerEngine 的其他 Issue

查看 NVIDIA/TransformerEngine 的全部 Issue

相似的 Issue

更多 Python Issue

把新 issue 发到你的邮箱

精选适合新手参与的 GitHub issue 摘要。