Prhub

#33703 [diffusion] Add SageAttention packed varlen path for MiniMax-H3

原始 PR 作者 niehen6174 合并时间 2026-08-06 01:19 文件变更 5 提交数 2 评论 1 代码增减 +82 / -14

执行摘要

H3 新增 SageAttention packed varlen 路径并解锁后端,提速约 5-7%

PR body 明确说明两点:一是 H3 attention 总是以 packed Q/K/V 调用 forward_varlen(cu_seqlens,约 3.8 万 tokens/step),而 sage_attn 后端只实现了 dense forward(),运行时直接 NotImplementedError;二是 SageAttention 暴露的是快速的 dense sageattn(),其 Triton sageattn_varlen 在 H3 长度序列上太慢(40.90s vs 本 PR 的 29.93s),因此需要把 packed 输入重新路由到 dense kernel。原 admission 和 pipeline 校验中拒绝 sage_attn 的理由是 'the current packed varlen path does not preserve model output',本 PR 用 fast path 解决了该保留问题。

值得精读。重点关注 _sage_packed 的双路径设计:如何用 dense kernel 模拟 varlen、如何用 _trailing_padding_used_len 识别 H3 特定打包布局、以及 padding 置零的技巧。对后续在 diffusion runtime 中接入其他 attention 后端有直接参考价值。

讨论亮点

本 PR 没有实质性的 review 评论,reviewer mickqian 直接 APPROVED,并通过 issue 评论 /tag-and-rerun-ci 重跑 CI。设计权衡主要在 PR body 中呈现:Triton varlen 路径太慢所以弃用,改用 dense sageattn 裁剪复用;trailing-padding fast path 明确假设 H3 的 bounds=(0, used, total) 布局。

实现拆解

  1. 新增 varlen 入口与打包路由:在 python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn.py 中新增 _trailing_padding_used_len() 纯函数,识别 H3 的 (0, used, total) 三元组打包布局(要求 start 为 0、used 小于 total、total 等于总 token 数、used 等于 max_seqlen);新增 forward_varlen() 作为 varlen 统一入口,优先使用 cu_seqlens_host 避免 GPU-CPU 同步,并对 QKV 做 contiguous() 预处理后交给 _sage_packed()
  2. 双模式 _sage_packed():fast path 在检测到 H3 单文档尾部 padding 时,将 [0, used) 切片提升 batch 维后调用现有 dense forward()(即 sageattn),padding 尾部用 torch.zeros_like 置零,依赖下游 masked 行保持 inactive;若 used 等于总 token 数则直接返回。兜底路径按 bounds 逐段调用 dense forward,与 SDPA varlen 后端的分段循环模式一致。forward() 本身保持不变,其他模型的 dense Sage 路径不受影响。
  3. 解锁 SAGE_ATTNpython/sglang/multimodal_gen/configs/models/dits/minimax_h3.py_supported_attention_backends 加入 AttentionBackendEnum.SAGE_ATTNpython/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/release_metadata.pyMiniMaxH3PartitionAdmissionStage.forward 删除对 sage_attn 的拒绝分支;python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.pyvalidate_server_args 同步删除拒绝逻辑,但 quality="high" 的严格 4xH200 部署校验仍然保留。
  4. 测试配套python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py 将原先 "sage_attn 应拒绝" 的断言改为 "sage_attn 应正常通过",同时保留 quality="high"/"ultra" 的拒绝路径覆盖。本 PR 未新增 forward_varlen 数值正确性的单元测试。
文件 模块 状态 重要度
python/sglang/multimodal_gen/runtime/layers/attention/backends/sage_attn.py 注意力后端 modified 8.03
python/sglang/multimodal_gen/runtime/pipelines_core/stages/model_specific_stages/minimax_h3/release_metadata.py 准入控制 modified 5.73
python/sglang/multimodal_gen/configs/pipeline_configs/minimax_h3.py 管线配置 modified 5.52
python/sglang/multimodal_gen/configs/models/dits/minimax_h3.py 模型配置 modified 4.39
python/sglang/multimodal_gen/test/unit/test_minimax_h3_admission.py 单元测试 modified 3.35

关键符号

_trailing_padding_used_len forward_varlen _sage_packed

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

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

正确性风险:fast path 强依赖 H3 的 bounds=(0, used, total) 且 used==max_seqlen 打包约定,若未来打包布局变化(多文档 packing 或非 64 对齐)会静默回退到分段循环,行为不同但不会崩溃;padding 行置零依赖下游 masked 计算,需确保所有下游算子都带 mask。性能风险:分段循环逐段 launch dense kernel,多文档 batch 下性能可能较差;forward_varlen 的 contiguous() 在非连续张量时带来额外拷贝;cu_seqlens_host 为 None 时 tolist() 引入 GPU-CPU 同步。回归风险:forward() 未改动,其他模型不受影响,但移除 admission 拒绝后所有用户都能以 SAGE_ATTN 启动 H3。测试缺口:缺少 forward_varlen 与 FA3/SDPA 的数值对比测试,PR 中 PSNR 27.7 dB 只是近似参考。

用户侧:H3 模型可直接选用 SageAttention,H200 上 denoise 约 5%、steady step 约 7% 提速,H20/L20/A100/4090 据 PR 描述更有竞争力。系统侧:新增一个注意力后端的 varlen 接入模式,为后续其他 packed 模型复用 Sage 铺路。团队侧:该模式(裁剪 fast path + 分段兜底)可作为 diffusion attention 后端的参考实现。

核心注意力路径变更 fast path 依赖 H3 打包约定 缺少 forward_varlen 数值正确性测试 移除 admission 拒绝

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论