Prhub

#29843 [trtllm_mha] Fuse cuda-graph metadata rebuild into one triton kernel

原始 PR 作者 pranjalssh 合并时间 2026-07-04 13:24 文件变更 9 提交数 7 评论 13 代码增减 +786 / -76

执行摘要

融合 TRTLLM MHA 图元数据为单 Triton 核

TRTLLM MHA 后端的 CUDA graph 元数据重建每次 replay 涉及约 25 个小 ATen 操作(index gather、floor_divide、cumsum 等),在某些 CPU 上纯调度耗时约 0.7-1.0 ms,且在每个解码步重复多次(draft-decode 步骤、target-verify、draft-extend)及每个 TP 秩上。由此产生的每秩 CPU jitter 导致 cudaGraphLaunch 跨秩偏差,该偏差被支付为首次自定义 all-reduce 中的旋转等待时间。

此 PR 值得精读,特别是融合 kernel 的设计模式(一次启动完成全部元数据更新、q_mode 抽象、SWA 翻译集成)。测试代码采用 monkeypatch 验证 graph 钩子行为,对类似 CUDA graph 优化有参考价值。建议团队关注 multi_layer_eagle 的测试补充。

讨论亮点

审查中主要关注点:

  • hybrid 后端转发:Qiaolin-Yu 指出混合注意力后端缺少 init_forward_metadata_in_graph,导致融合内核在混合模型中不生效。作者立即添加转发方法。
  • frozen_kv_mtp 测试:Qiaolin-Yu 询问是否经过测试,作者添加了专用的 runner 测试用例。
  • multi_layer_eagle 测试:Qiaolin-Yu 提出类似问题,作者询问推荐测试方法,目前无明确结论。
  • SWA 支持:ch-wan 质疑 out_graph 中删除 SWA 代码后支持是否保留,merrymercy 解释 SWA 翻译已通过 _apply_cuda_graph_metadata 中的融合内核处理。

实现拆解

  1. 新增融合 Triton 内核:在 trtllm_mha_graph_metadata.py 中定义 update_trtllm_mha_graph_metadata_kernel,利用每个 batch 行一个 program 的并行模式,一次计算所有元数据。支持三种 q_mode(预设/cumsum/strided)、SWA 翻译和 -1 sentinel 保护。
  2. 后端集成替换:在 trtllm_mha_backend.py 中导入新内核,重写 _apply_cuda_graph_metadata,将原先约 25 个操作替换为单次内核调用。同时预处理 SWA 映射(_swa_full_to_swa_mapping)供内核使用。init_forward_metadata_in_graph 现在调用融合内核作为 graph 录制的一部分。
  3. 推测解码运行器扩展:在 eagle_draft_extend_cuda_graph_runner.pyfrozen_kv_mtp_cuda_graph_runner.pymulti_layer_eagle_draft_extend_cuda_graph_runner.pycapture_one_shape 中插入 init_forward_metadata_in_graph 调用。draft-extend 路径使用静态捕获的 q_stride 避免运行时 .item() 调用。
  4. 混合注意力后端修复hybrid_attn_backend.pyhybrid_linear_attn_backend.py 原未转发 in-graph 钩子,本 PR 添加 init_forward_metadata_in_graph 方法以委托给内部后端。
  5. 测试配套:新增 test_trtllm_mha_graph_metadata.py(432 行),包含基于纯 ATen 参考的正确性测试(覆盖 batch 组合、q_mode、SWA 开关、零填充尾部),以及钩子分发测试和 graph 录制重放测试。现有 test_trtllm_mha.py 增加 frozen_kv_mtp 的 runner 测试。
文件 模块 状态 重要度
test/registered/attention/test_trtllm_mha_graph_metadata.py MHA 测试 added 8.14
python/sglang/srt/layers/attention/trtllm_mha_backend.py MHA 后端 modified 7.89
python/sglang/srt/layers/attention/triton_ops/trtllm_mha_graph_metadata.py 元数据核 added 6.95
python/sglang/srt/layers/attention/hybrid_attn_backend.py 混合注意力 modified 5.77
test/registered/attention/unittests/dense/test_trtllm_mha.py MHA 测试 modified 5.61
python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py 线性注意力 modified 5.49
python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py 推测解码 modified 5.24
python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py 推测解码 modified 4.54
python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py 推测解码 modified 4.54

关键符号

update_trtllm_mha_graph_metadata_kernel update_trtllm_mha_graph_metadata init_forward_metadata_in_graph _apply_cuda_graph_metadata _build_cuda_graph_metadata

分析完成后,这里会展示 LLM 生成的相对完整源码片段和详细注释。

评论区精华

混合注意力后端缺少 init_forward_metadata_in_graph 转发 正确性

Qiaolin-Yu 指出 hybrid_attn_backend.py 和 hybrid_linear_attn_backend.py 未实现 init_forward_metadata_in_graph,导致融合内核在混合模型中不执行。

结论:pranjalssh 同意并添加了对应转发实现。 · 已解决

Frozen-KV MTP 图录制需要测试验证 测试

Qiaolin-Yu 询问是否对此修改进行测试。

结论:pranjalssh 添加了 test_runner_mode_frozen_kv_mtp_cuda_graph_runner_cases。 · 已解决

Multi-layer eagle draft extend 图录制测试方法 测试

Qiaolin-Yu 对类似修改提出相同问题,pranjalssh 询问推荐测试方法。

结论:暂未解决,需后续跟进。 · unresolved

融合内核后 SWA 支持是否保留 正确性

ch-wan 问 init_forward_metadata_out_graph 中删除了 SWA 相关代码,SWA 如何支持?

结论:merrymercy 解释 SWA 翻译已在融合内核的 _apply_cuda_graph_metadata 中处理。 · 已解决

风险与影响

  1. SWA 映射一致性:SWA 依赖在构造时捕获的 full_to_swa_mapping 索引映射,若运行时池发生变化但映射未更新,可能导致页表翻译错误。但 SWA 池通常稳定,风险较低。
  2. multi_layer_eagle 测试空白:该路径无直接测试覆盖,可能引入回归。但已有 EAGLE 链式测试,类似逻辑。
  3. Blackwell 独占:新内核仅 sm100 可用,其他 GPU 仍使用原有 ATen 路径,功能无影响,但性能优化不覆盖。
  4. 混合后端转发完整性:虽然添加了转发,但若新增后端类型未正确转发,in-graph 钩子可能静默缺失。但现有测试覆盖了 hybrid 后端的基本路径。
  • 用户:Blackwell GPU 用户在多 TP 场景下应感受到降低的延迟抖动和可能的吞吐提升;其他 GPU 用户无功能变化。
  • 系统:减少 host-device 交互,降低 CPU 调度负载,使 cudaGraphLaunch 跨秩更同步。
  • 团队:需维护新的 Triton kernel,并确保与未来后端演进兼容。
SWA 映射依赖构造时间快照 multi_layer_eagle 测试未覆盖 kernel 仅 Blackwell(sm100) 可用 混合后端转发可能遗漏

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论