Prhub

#34217 [misc] Pass FP8 scales in FlashInfer SWA prefill, autotune fp8 on SM120, and tighten `is_image_understandable_model`

原始 PR 作者 hnyls2002 合并时间 2026-08-10 14:45 文件变更 12 提交数 8 评论 3 代码增减 +117 / -244

执行摘要

修复 FlashInfer SWA 漏传 FP8 scale,SM120 开启 fp8 autotune

PR body 明确说明三处动机:FlashInferAttnBackend 有三个 paged-KV wrapper 调用,其中两处传了 k_scale/v_scale,滑动窗口的 forward_return_lse 漏传,导致「a model combining an FP8 KV cache with a sliding window reads the cache as if it were unscaled」;SM120 满足 is_blackwell_supported(),resolve_mxfp8_dense_gemm_backend 会把它路由到与 SM100 相同的可调优 MXFP8 dense GEMM,「without this the kernel always runs at tactic=-1 there」;is_image_understandable_model 的 hasattr 问的是属性是否存在,而意图是视觉塔是否存在,deepseek_ocr、deepseekvl2、unlimited_ocr、janus_pro 都把 vision_config 声明为 None 默认值,此类配置会被误报为 image-capable 并在 warmup 时用图像探测。作者强调这是正确性收紧,不是已观测 bug。

值得精读 flashinfer_backend.py 的修复模式——同一 wrapper 对象多处调用的一致性排查方法,以及 flashinfer_autotune.py 中「按硬件能力判断而非架构白名单」的思路。对在 Blackwell 上部署 FP8 KV cache + 滑动窗口模型的团队有直接意义;MLX 注释清理部分可略读。建议后续为 forward_return_lse 的 scale 传递补充一条针对 FP8 + SWA 的回归测试,并把 test_fused_fp8_kv_write.py 注册进 CI。

讨论亮点

本 PR 没有 code review 评论,审查意见主要通过 issue 内命令与 commit 演进体现:

  • 作者通过 /rerun-test 指定了三组测试:test_flashinfer.py(4-gpu-b200、1-gpu-h100)、test_model_config.py(ubuntu-latest),三组均通过;test_fused_fp8_kv_write.py 被 bot 以「No register_cuda_ci(...) or register_cpu_ci() found」拒绝执行,暴露该文件未注册 CI。
  • commit 97c1b7f 标题为「exempt only sm120 from mxfp8 autotune skip」,说明最初思路是在环境变量 skip 名单上豁免 SM120;最终改为直接把 SM120 纳入 autotune 条件,语义更明确且不依赖环境变量。bdd1b1f 又把注释更新为点名 resolve_mxfp8_dense_gemm_backend 实际的 SM120 路由行为。
  • 最后一个 commit 是 merge main 并解决 flashinfer_autotune.py 冲突,说明该文件近期有并行改动(与 #33962 相关)。

实现拆解

  1. 修复 FlashInfer SWA prefill 的 FP8 scale 漏传(核心):在 python/sglang/srt/layers/attention/flashinfer_backend.pyforward_extend 中,给 prefill_wrapper_paged.forward_return_lse 补充 k_scale/v_scale 关键字参数,对齐同函数另外两处 paged 调用(prefill_wrapper_paged.forwarddecode_wrapper.forward)。两处 ragged 调用保持省略 scale 不变——它们读的是当前 batch 的 k/v 而非 paged cache。非 FP8 场景下 k_scale_float/v_scale_float 默认 None,与改动前完全等价,回归面仅限 FP8 + SWA 组合。
  2. 解锁 SM120 的 fp8 autotune:在 python/sglang/srt/model_executor/runner/flashinfer_autotune.pyshould_run_flashinfer_autotune 中,fp8 分支由 is_sm100_supported() 扩展为 is_sm100_supported() or is_sm120_supported()。原因是 SM120 同样满足 is_blackwell_supported(),会被 resolve_mxfp8_dense_gemm_backend 路由到可调优的 FlashInfer CUTLASS MXFP8 dense GEMM,但 autotune 条件此前不含 SM120,内核始终以 tactic=-1 运行。PR 初始 commit 曾尝试针对 mxfp8_gemm 的环境变量方案,后因 #33962 已清理 FLASHINFER_AUTOTUNE_WORKAROUND_SKIPS 而放弃,最终只保留条件判断。
  3. 收紧多模态能力判据:在 python/sglang/srt/configs/model_config.pyModelConfig.__init__ 中,is_image_understandable_modelhasattr(self.hf_config, "vision_config") 改为 getattr(self.hf_config, "vision_config", None) is not None。理由是 deepseek_ocrdeepseekvl2unlimited_ocrjanus_pro 等配置类把 vision_config 声明为类属性且默认 None——属性存在不等于视觉塔存在;此前这类配置会被误判为 image-capable 并在 warmup 时用图像探测。当前无模型实际触发,属防御性收紧。
  4. 精简 MLX 后端及相关模块注释hardware_backend/mlx/ 下的 sampling.pymodel_runner.pytp_worker.pykv_cache/attention_kv_cache.pykv_cache/attention_wrapper.pyscheduler_mixin.py,以及 arg_groups/overrides.pyconstrained/xgrammar_backend.pyserver_args.py 中的 MLX 相关块,删除重复叙述和超长背景说明(如 #25804 的历史、gpt-oss 层数示例、~100x 量化对比等)。PR 声明所有涉及文件去掉 docstring 与注释后与 main AST 一致,纯注释变更。
  5. 验证配套:未新增测试文件,通过 /rerun-test4-gpu-b2001-gpu-h100 上运行 test/registered/attention/unittests/swa/test_flashinfer.py,在 ubuntu-latest 上运行 test/registered/unit/configs/test_model_config.py;尝试运行 test/registered/attention/test_fused_fp8_kv_write.py 时被 CI bot 拒绝——该文件未注册 register_cuda_ci/register_cpu_ci
文件 模块 状态 重要度
python/sglang/srt/layers/attention/flashinfer_backend.py 注意力后端 modified 5.68
python/sglang/srt/model_executor/runner/flashinfer_autotune.py 自动调优 modified 5.84
python/sglang/srt/configs/model_config.py 模型配置 modified 5.11
python/sglang/srt/hardware_backend/mlx/sampling.py 采样器 modified 5.7
python/sglang/srt/hardware_backend/mlx/model_runner.py 模型运行器 modified 5.4

关键符号

forward_extend should_run_flashinfer_autotune

关键源码片段

python/sglang/srt/layers/attention/flashinfer_backend.py core-logic

核心修复:SWA chunked-prefill 的 paged 读取补传 FP8 反量化 scale,是唯一改变运行行为的注意力路径变更。

# python/sglang/srt/layers/attention/flashinfer_backend.py · forward_extend 内三处
# paged-KV 调用(依 PR 描述重建的示意片段,非原始窗口):
# FlashInferAttnBackend 对 paged cache 的读取有三处,其中两处已传 KV 反量化
# scale,只有 SWA chunked-prefill 的 forward_return_lse 漏传,导致 FP8 KV cache
# + 滑动窗口组合按未缩放缓存读取。# 1) 常规 paged prefill:k_scale / v_scale 已传递 ✓
prefill_wrapper_paged.forward(
    q, k, v,
    ...,
    k_scale=k_scale_float,
    v_scale=v_scale_float,
)# 2) SWA chunked-prefill(forward_return_lse):此前漏传 scale,本次补齐 ✓
out, lse = prefill_wrapper_paged.forward_return_lse(
    q, k, v,
    ...,
    k_scale=k_scale_float, # ← 本次新增
    v_scale=v_scale_float, # ← 本次新增
)# 3) decode:k_scale / v_scale 已传递 ✓
decode_wrapper.forward(
    q, k, v,
    ...,
    k_scale=k_scale_float,
    v_scale=v_scale_float,
)# 两处 ragged 调用(读当前 batch 的 k/v,而非 paged cache)保持省略 scale。
# 非 FP8 场景下 k_scale_float / v_scale_float 默认 None,与改动前完全等价。
python/sglang/srt/model_executor/runner/flashinfer_autotune.py core-logic

SM120 纳入 fp8 autotune 条件,消除该架构上 MXFP8 dense GEMM 固定 tactic=-1 的性能损失。

# python/sglang/srt/model_executor/runner/flashinfer_autotune.py
# should_run_flashinfer_autotune 的 fp8 分支:SM120 也满足 is_blackwell_supported(),
# 会被 resolve_mxfp8_dense_gemm_backend 路由到与 SM100 相同的 FlashInfer CUTLASS
# MXFP8 dense GEMM(可调优内核)。若不把 SM120 纳入条件,该内核在 SM120 上
# 永远以 tactic=-1 运行,性能受损。from sglang.srt.utils import is_sm100_supported, is_sm120_supportedmodel_uses_modelopt_fp8 = model_quantization in (
    "modelopt",
    "modelopt_fp8",
    "modelopt_mixed",
)# SM120 满足 is_blackwell_supported(),所以 resolve_mxfp8_dense_gemm_backend
# 会把它送到与 SM100 同一个可调优的 FlashInfer CUTLASS MXFP8 dense GEMM;
# 没有这个条件时内核在 SM120 上始终以 tactic=-1 运行。
fp8_gemm_needs_autotune = get_fp8_gemm_runner_backend().is_flashinfer_cutlass() or (
    model_uses_modelopt_fp8 and (is_sm100_supported() or is_sm120_supported())
)
python/sglang/srt/configs/model_config.py data-contract

多模态能力判据从属性存在性收紧为 vision_config 非 None,影响模型能力判定与 warmup 行为。

# python/sglang/srt/configs/model_config.py · ModelConfig.__init__
# 关键在「视觉塔是否存在」而非「属性是否存在」:deepseek_ocr、deepseekvl2、
# unlimited_ocr、janus_pro 等配置类把 vision_config 声明为类属性且默认 None,
# hasattr 判断会把它们误判为 image-capable,导致 warmup 用图像探测;
# getattr 取值判 None 后,此类配置如实落入纯文本路径。self.is_image_understandable_model = (
    enable_multimodal
    and not self.is_lm_only
    and getattr(self.hf_config, "vision_config", None) is not None
)

评论区精华

/rerun-test 验证 SWA FlashInfer 与模型配置测试 测试

作者通过 /rerun-test 指令在 4-gpu-b200 和 1-gpu-h100 上运行 test/registered/attention/unittests/swa/test_flashinfer.py,在 ubuntu-latest 上运行 test/registered/unit/configs/test_model_config.py,三组均返回 ✅。

结论:SWA FP8 prefill 修复与 is_image_understandable_model 收紧在目标硬件和 CPU 侧均通过验证。 · 已解决

test_fused_fp8_kv_write.py 未注册 CI 被拒绝执行 测试

CI bot 回复:No register_cuda_ci(runner_config=...) or register_cpu_ci() found in test/registered/attention/test_fused_fp8_kv_write.py,该文件可能不是注册的 CI 测试。

结论:该 FP8 KV 写入测试未纳入 CI 注册体系,无法通过 /rerun-test 运行,需另行注册。 · unresolved

风险与影响

  1. flashinfer_backend.py(核心修复):仅新增两个默认 None 的关键字参数,非 FP8 场景行为与改动前完全一致;FP8 + SWA 场景从错误读取变为正确缩放,方向安全。但该修复没有配套新增单测,依赖已有 test_flashinfer.py 的覆盖,若未来重构 forward_return_lse 的返回值契约(out, lse)需注意参数位置不被破坏。
  2. flashinfer_autotune.py:SM120 首次启动会增加 autotune 搜索耗时(tactic 扫描),换来脱离固定 tactic=-1 的长期性能收益;autotune 缓存路径按 sm_major_minor 归档,与 SM100 缓存天然隔离,无串扰风险。
  3. model_config.py:判据收紧后,若未来某模型在 __init__ 之后才惰性初始化 vision_config,该谓词(在 __init__ 内求值)会先读到 None 而误判为纯文本;当前所有相关配置类都在 __init__ 内赋真实对象,风险低。
  4. MLX 注释清理:AST 一致无行为风险,但删除了「Un-gate this together with the window-aware shared SWA pool」这类指引性上下文,未来开发者在 init_cache_pools 等处可能丢失这段设计意图。
  5. 测试缺口:核心 FP8 scale 修复没有直接对应新增测试,test_fused_fp8_kv_write.py 又未注册 CI,回归验证主要靠既有的 SWA FlashInfer 测试。

对用户:修复 SM100/SM120 上 FP8 KV cache 与滑动窗口模型组合的潜在静默精度问题;SM120 用户获得 MXFP8 dense GEMM 的 autotune 收益,不再固定运行在 tactic=-1。对系统:多模态模型启动 warmup 的判据更准确,避免对纯文本配置误用图像探测。对团队:MLX 后端约 244 行注释净删除,降低文档维护噪音;同时揭示 test_fused_fp8_kv_write.py 未注册 CI 的流程问题。整体影响面集中于注意力后端与模型配置两个模块,范围可控。

核心注意力路径变更 SM120 autotune 启动开销 模型能力判据语义收紧 FP8 修复缺少直接单测 注释清理无测试覆盖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论