Prhub

#46114 [ROCm][Bugfix] Fix chunk alignment when using context parallelism with TRITON_MLA

原始 PR 作者 micah-wil 合并时间 2026-06-25 12:07 文件变更 2 提交数 9 评论 3 代码增减 +53 / -39

执行摘要

修复 ROCm MLA context parallelism 的 chunk 对齐问题

There is a bug in mla_attention when using context parallelism on ROCm. max_context_chunk is being aligned properly on CUDA because of the self.aot_schedule path (which is CUDA-only). The chunk misalignment causes zero accuracy in test_context_parallel.py on ROCm with dcp_size=4 using TRITON_MLA.

值得精读,特别是对于涉及跨平台条件编译的项目。该 PR 演示了如何通过移除不必要的平台特定门控来修复 bug,并展示了测试重构以提升可靠性的典型手法。

讨论亮点

Review 中 LucasWilkinson 提出质疑:为什么不直接使用 round_down(max_context_chunk, self.kv_cache_spec.block_size) 而是引入 lcm?作者 micah-wil 反问能否移除 self.aot_schedule 门控。Lucas 指出该门控已是历史遗留(最初因为 self.kv_cache_spec.block_size 需要从 runner 获取,现已可直接使用)。最终共识:移除 aot_schedule 条件,让对齐逻辑无条件执行,简化代码。

实现拆解

  1. 移除平台条件门控:在 MLACommonAttention.__init__ 中,原代码仅当 self.aot_schedule(CUDA)时才设置 self.page_size,导致 ROCm 上 page_size 未初始化。现改为无条件设置 self.page_size = self.kv_cache_spec.block_size
  2. 统一 chunk 对齐逻辑:在 MLACommonAttention.build 方法中,原代码的 round_down 对齐操作仅在 self.aot_schedule 条件下执行,现移除该条件,使 max_context_chunk 始终按 page_size 对齐,确保所有 GPU 平台一致行为。
  3. 测试配套重构:在 tests/distributed/test_context_parallel.py 中,根据 current_platform.is_rocm() 分流模型配置,ROCm 下仅测试 dcp_multipliers=[1](避免 CUDA 专属后端)。同时将评估方法从自定义 evaluate_gsm8k 切换为 lm_eval.simple_evaluate,使用 local-completions 接口,提高结果稳定性。
文件 模块 状态 重要度
vllm/model_executor/layers/attention/mla_attention.py 注意力层 modified 6.54
tests/distributed/test_context_parallel.py 并行测试 modified 5.52

关键符号

MLACommonAttention.__init__ MLACommonAttention.build

关键源码片段

vllm/model_executor/layers/attention/mla_attention.py core-logic

核心修复文件:移除平台条件门控,使 page_size 初始化和 chunk 对齐逻辑在所有 GPU 平台上统一执行。

# 在 __init__ 中:移除平台条件,统一设置 page_size
self.page_size = self.kv_cache_spec.block_size# 在 build() 方法中处理 chunked context 对齐:
if max_context_len_cpu > 0:
    max_context_chunk = (
        self.chunked_prefill_workspace_size // num_prefills_with_context_cpu
    )
    # 之前只在 aot_schedule (CUDA) 下对齐,现在无条件执行 round_down
    max_context_chunk = round_down(max_context_chunk, self.page_size)
    assert max_context_chunk > 0
    num_chunks = cdiv(max_context_len_cpu, max_context_chunk)

评论区精华

简化 chunk 对齐逻辑,移除平台条件门控 设计

LucasWilkinson 询问为何需要 lcm,建议直接 round_down 并移除 aot_schedule 门控;micah-wil 确认可行并询问能否彻底移除 aot_schedule;LucasWilkinson 解释该条件已是历史遗留,可直接移除。

结论:采用直接无条件 round_down 方案,删除平台门控,简化代码。 · 已解决

风险与影响

该 PR 改动集中在注意力层核心,但改动量小且逻辑清晰:将原本 CUDA 专属的对齐行为推广到所有平台。潜在风险是如果未来某些平台需要不同的对齐策略,该统一逻辑可能不适用,但当前所有 GPU 平台均使用相同的 block_size 语义,风险较低。测试已覆盖 ROCm 下的 context parallelism 场景,回归风险较小。

直接影响使用 ROCm + TRITON_MLA + context parallelism (dcp > 1) 的用户,修复了此前零准确率的问题。对 CUDA 平台无功能变化(对齐逻辑与之前一致)。测试改用 lm_eval 带来更稳定的 CI 结果,但引入了新的依赖 (lm_eval)。

核心路径变更 平台条件依赖移除 测试增强覆盖

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论