Prhub

#33655 [diffusion] Prefer cuDNN SDPA over FA4 for dense attention on sm_100 (B200)

原始 PR 作者 BBuf 合并时间 2026-08-06 08:49 文件变更 8 提交数 5 评论 4 代码增减 +180 / -31

执行摘要

B200 扩散模型默认切 cuDNN SDPA,比 FA4 快 1.3-1.5 倍

PR body 指出 NVlabs Sana sol-engine 分支在 sm_100 上用 cuDNN SDPA 替换 flash-attn 2.8.3 varlen 拿到 1.81x 端到端收益,但该结论基于 FA2 baseline,不直接适用于 sglang:sglang 在 sm_100 上默认分发的是 vendored FA4 CuTe DSL 内核。作者在 B200(torch 2.11.0+cu130、cuDNN 9.19)上重新实测,发现 FA4 虽已快于 FA2 体系,cuDNN 9.19 SDPA 仍比 FA4 快 1.24-1.98x(跨 Wan2.2、LingBot-World、MiniMax-H3、Qwen-Image、FLUX 全部真实形状),因此现有 sm_100 默认分发是次优的,需要修正默认值。

值得精读。该 PR 是一个「以数据驱动改变默认值」的优秀范例:先修正外部结论(FA2 baseline → FA4 baseline),再用 11 个真实 shape 的 microbenchmark 和 9 轮端到端 A/B 支撑决策;同时通过 fail-safe 链、NVFP4 数值保护和单测把变更风险压到最小。值得关注的设计决策包括:sm_100 专用 gate 与 Hopper 启发式并存、_cudnn_failed 按层锁存避免重复探测、return_softmax_lse 强制走 FA 以兼容 ring attention、以及加载期 context manager 实现量化默认后端的局部覆盖。

讨论亮点

PR 本身没有 review 评论,但合并后 issue 评论区出现了一次完整的 H100 回归误报调查,过程很有价值:

  • mickqian 首先报告 flux_image_t2i_2_gpus 在 H100 2-GPU 套件上 Denoise Step 从 73.6 ms 基线涨到 420-656 ms,怀疑「要么 sm_100 gate 泄漏到 H100,要么新选择路径每次 step 有 probe 开销」,并给出基于 #33725 CI 时间线(00:56 失败,本 PR 00:49 落地)的归因证据。
  • 随后他阅读 diff 后做了机制分析:失败日志中 cuDNN SDPA failed ... falling back to FlashAttention 警告出现 0 次,_is_sm100 在 H100 上读取正确,说明运行时 fallback 路径不是原因;而回归是每个 step 恒定的 6-9x 而非 step 0 单次尖峰,指向 load-time 默认 resolution 的副作用。
  • 关键转折:mickqian 更正称自己的归因是错的——失败 job 的 checkout 时间是 00:18 UTC,早于本 PR 落地 31 分钟,运行不可能包含本变更;离线 A/B(ba12a16^ vs ba12a16,H200 ABAB×2)两臂一致健康(约 82 ms/step)。
  • 最终结论:同一 merge commit 在 rerun(attempt 2)上全绿,失败原因是 runner 状态(当天该 runner 有 incomplete-weight-cache 和 OOM 事故),不存在代码回归。

这次讨论展示了如何用 checkout 时间、日志出现频率、per-step 常量因子等证据链做回归归因,并最终以 rerun 证伪。

实现拆解

该 PR 的核心是一次「仅针对 sm_100 默认值」的分发调整,不触碰任何 kernel 代码,配套了量化数值保护和测试。实现按以下步骤推进:

  1. 默认后端解析入口改造(python/sglang/multimodal_gen/runtime/platforms/cuda.py:在 CudaPlatform.get_attn_backend_cls_strselected_backend is None 分支里,当自动解析结果为 AttentionBackendEnum.FAcls.is_blackwell()(compute capability 10.x)时,先调用 _resolve_flash_attention_backend_cls_str 确认 FA 实际可用(head size、dtype 等 guard 不变);若 FA 可用则返回 _DYNAMIC_CUDNN_SDPA_BACKEND_CLS_STR,若 FA 不可用(会落到 Torch SDPA)则保持原 Torch SDPA 路径。显式 --attention-backend 选择不受影响。

  2. 运行时双实现调度(python/sglang/multimodal_gen/runtime/layers/attention/backends/sdpa.pyDynamicCudnnSDPAImpl 新增 _is_sm100get_device_capability()[0] == 10)和 _cudnn_failed 状态锁存;_use_cudnn_sdpa 在 sm_100 上对非 causal、CUDA、fp16/bf16、Sq == Skv 的 dense 形状直接返回 True(同时支持 cross-attn,Skv 为文本长度),非 sm_100 保留原有 Hopper 启发式(D=64、S=1024、B>=4);forward 增加 **kwargs,当调用方要求 return_softmax_lse(如 ring attention)时强制走 FA,并对 cuDNN 抛出的 RuntimeError 做一次 warning 后永久钉回 FA,避免每个 step 重复探测失败 kernel。

  3. ModelOpt NVFP4 数值保护(transformer_load_utils.py + transformer_loader.pyTransformerQuantLoadSpec 新增 is_modelopt_fp4 property;transformer_loader 新增 _default_quantized_attention_backend,在 Blackwell + ModelOpt FP4 且没有任何显式/全局/组件级 attention 后端时返回 AttentionBackendEnum.FA,并在 load_customized 中通过 component_attn_backend_context_manager 把该默认包在 maybe_load_fsdp_model 构造期周围(无默认时用 nullcontext)。这是为了防止新的 cuDNN 默认路径改变 NVFP4 量化模型的数值行为。

  4. 测试配套test_cuda_attention_backend.py 新增 test_default_backend_prefers_dynamic_cudnn_sdpa_on_blackwell,断言 Blackwell 默认解析出 DynamicCudnnSDPABackend 类名;test_transformer_quant.py 新增两个用例分别验证 FP4 默认走 FA、显式后端不被覆盖;test_glm_image_ar.py 引入 _FakeBatchResponse 以匹配 GLM-Image AR 批式返回新契约,test_glm_image_multi_output.py 补充 extra={} 字段并适配 generate_prior_tokens 三元组返回。另在 PR body 中提供了 B200 上的 microbenchmark、9 轮端到端 A/B 和 PSNR/SSIM 质量对比数据作为验证。

文件 模块 状态 重要度
python/sglang/multimodal_gen/runtime/layers/attention/backends/sdpa.py 注意力后端 modified 6.76
python/sglang/multimodal_gen/runtime/platforms/cuda.py 平台选择器 modified 5.94
python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py 模型加载 modified 7.13
python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py 量化加载 modified 5.09
python/sglang/multimodal_gen/test/unit/test_transformer_quant.py 量化测试 modified 5.53
python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py 后端测试 modified 4.75
python/sglang/multimodal_gen/test/unit/test_glm_image_ar.py AR 测试 modified 5.96
python/sglang/multimodal_gen/test/unit/test_glm_image_multi_output.py 多输出测试 modified 3.32

关键符号

_default_quantized_attention_backend TransformerQuantLoadSpec.is_modelopt_fp4 CudaPlatform.get_attn_backend_cls_str DynamicCudnnSDPAImpl.__init__ DynamicCudnnSDPAImpl._use_cudnn_sdpa DynamicCudnnSDPAImpl.forward

关键源码片段

python/sglang/multimodal_gen/runtime/layers/attention/backends/sdpa.py core-logic

sm_100 默认路径的实际运行时实现:新增 _is_sm100 判断、_cudnn_failed 锁存、sm_100 dense 形状直通 cuDNN SDPA,并处理 return_softmax_lse 与 cuDNN 异常兜底。

# DynamicCudnnSDPAImpl 是 sm_100 上默认 attention 路径的运行时实现:
# cuDNN SDPA 为主、FA4 为兜底,并按 layer 级永久锁存失败状态。
class DynamicCudnnSDPAImpl(SDPAImpl):
    def __init__(self, num_heads, head_size, causal, softmax_scale, num_kv_heads=None, prefix="", **extra_impl_args):
        from sglang.multimodal_gen.runtime.layers.attention.backends.flash_attn import (
            FlashAttentionImpl, set_fa_ver,
        )
​
        self.causal = causal
        self.head_size = head_size
        # sm_100 特判:compute capability 主版本等于 10,Hopper 与 sm_120 不受影响
        self._is_sm100 = (
            torch.cuda.is_available() and torch.cuda.get_device_capability()[0] == 10
        )
        # cuDNN SDPA 一旦在某层抛错,永久钉死 FA 路径,避免每个 step 重复探测失败 kernel
        self._cudnn_failed = False
        if torch.cuda.is_available() and torch.cuda.get_device_capability()[0] >= 10:
            set_fa_ver(4)
        self.cudnn_impl = CudnnSDPAImpl(
            num_heads=num_heads, head_size=head_size, causal=causal,
            softmax_scale=softmax_scale, num_kv_heads=num_kv_heads,
            prefix=f"{prefix}.cudnn", **extra_impl_args,
        )
        self.fa_impl = FlashAttentionImpl(
            num_heads=num_heads, head_size=head_size, causal=causal,
            softmax_scale=softmax_scale, num_kv_heads=num_kv_heads,
            prefix=f"{prefix}.fa", **extra_impl_args,
        )
​
    def _use_cudnn_sdpa(self, query, key, value):
        # causal 或已锁存失败时直接走 FA;其余 guard 保持原语义
        if self.causal or self._cudnn_failed:
            return False
        if query.device.type != "cuda":
            return False
        if query.dtype not in (torch.float16, torch.bfloat16):
            return False
        if query.shape[2] != key.shape[2]:
            return False
        if self._is_sm100:
            # B200/sm_100 实测 cuDNN SDPA 比 FA4 CuTe 快 1.25-1.5x;
            # 覆盖 dense self-attn(Sq == Skv,最长 506K)与 cross-attn(Skv = 文本长度)
            return True
        # 非 sm_100 保留原有 Hopper 启发式:仅 D=64、S=1024、B>=4 走 cuDNN
        if query.shape[1] != key.shape[1]:
            return False
        return query.shape[-1] == 64 and query.shape[1] == 1024 and query.shape[0] >= 4
​
    def forward(self, query, key, value, attn_metadata, **kwargs):
        # return_softmax_lse 只有 FA 实现支持(如 ring attention),必须强制走 FA
        if not kwargs.get("return_softmax_lse") and self._use_cudnn_sdpa(query, key, value):
            try:
                return self.cudnn_impl.forward(query, key, value, attn_metadata)
            except RuntimeError as e:
                # cuDNN 可能对个别 shape 报 “No available kernel”;记录一次后本层永久回退 FA
                logger.warning(
                    "cuDNN SDPA failed (%s); falling back to FlashAttention for %s.",
                    e, type(self).__name__,
                )
                self._cudnn_failed = True
        return self.fa_impl.forward(query, key, value, attn_metadata, **kwargs)
python/sglang/multimodal_gen/runtime/loader/component_loaders/transformer_loader.py core-logic

新增 _default_quantized_attention_backend 并用 attention backend context manager 包裹 FSDP 模型构造,为 ModelOpt NVFP4 保留 FA4 数值路径。

# 在 Blackwell 上,ModelOpt NVFP4 模型默认保持 FA4,避免 cuDNN SDPA
# 改变 bf16 数值路径、破坏量化的稳定输出;用户显式选择则优先。
def _default_quantized_attention_backend(
    quant_spec: TransformerQuantLoadSpec, server_args: ServerArgs
) -> AttentionBackendEnum | None:
    if not current_platform.is_blackwell() or not quant_spec.is_modelopt_fp4:
        return None
    if (
        get_global_forced_attn_backend() is not None
        or get_component_forced_attn_backend() is not None
        or server_args.attention_backend is not None
    ):
        return None
    return AttentionBackendEnum.FA# 加载 transformer 组件时,把量化默认后端应用到“模型构造期”:
# attention 实现在构造时解析,因此用 context manager 包住 FSDP init + load。
quantized_attn_backend = _default_quantized_attention_backend(
    quant_spec, component_server_args
)
if quantized_attn_backend is not None:
    logger.info(
        "Using %s attention for ModelOpt NVFP4 to preserve output precision",
        quantized_attn_backend.name.lower(),
    )
attn_backend_context = (
    component_attn_backend_context_manager(
        quantized_attn_backend, component_name=component_name
    )
    if quantized_attn_backend is not None
    else nullcontext()
)
with attn_backend_context:
    model = maybe_load_fsdp_model(
        model_cls=model_cls,
        init_params=init_params,
        weight_dir_list=safetensors_list,
        device=local_torch_device,
        hsdp_replicate_dim=server_args.hsdp_replicate_dim,
        hsdp_shard_dim=server_args.hsdp_shard_dim,
        cpu_offload=component_server_args.dit_cpu_offload,
        pin_cpu_memory=component_server_args.pin_cpu_memory,
        fsdp_inference=component_server_args.use_fsdp_inference,
        param_dtype=quant_spec.param_dtype,
        reduce_dtype=torch.float32,
        output_dtype=None,
        strict=False,
        weight_load_plan=weight_load_plan,
    )

评论区精华

H100 flux_image_t2i_2_gpus 回归误报 性能

mickqian 报告 H100 2-GPU 套件上 flux Denoise Step 从 73.6 ms 基线涨到 420-656 ms,基于 #33725 CI 时间线推测与本 PR 的 sm_100 选择逻辑有关,怀疑 gate 泄漏或每步 probe 开销。

结论:回归报告被作者自己推翻:失败 job 的 checkout 时间早于本 PR 落地 31 分钟,不可能包含本变更;离线 A/B 两臂一致,最终 rerun 全绿,确认为 runner 状态问题。 · 已解决

回归机制分析:fallback 路径被排除 正确性

mickqian 阅读 diff 后指出失败日志中无 “cuDNN SDPA failed” 警告、_is_sm100 对 H100 读取正确,且回归是每个 step 恒定 6-9x 而非首步尖峰,因此指向 load-time 默认 resolution 而非 sdpa.py 运行时分发。

结论:该机制分析本身有效,但因 checkout 时间判断错误而指向了错误嫌疑对象;最终证伪后,分析思路仍值得保留。 · 已解决

最终归因:runner 状态而非代码回归 other

mickqian 最终确认同一 merge commit 在 rerun(attempt 2)上全绿,失败起因是 runner 当天的 incomplete-weight-cache 与 OOM 事故;同时排除了 sgl-kernel 0.4.6 与 kernel dispatch 统一两个嫌疑。

结论:本 PR 与嫌疑 commit 均无代码回归,case closed。 · 已解决

风险与影响

  1. 平台泄漏风险(已证伪但需留意):默认切换由 is_blackwell()(cc 10.x)和 _is_sm100 双重守卫,sm_120、Hopper 及更早平台不受影响;H100 回归误报最终被证实为 runner 状态问题,而非代码路径泄漏。
  2. 运行时 fallback 依赖异常类型DynamicCudnnSDPAImpl.forward 只捕获 RuntimeError 并永久钉回 FA。若 cuDNN 对某 shape 返回非 RuntimeError 的异常(如 torch.cuda.OutOfMemoryError 继承自 RuntimeError,会被错误地当作 kernel 不支持而静默降级),可能掩盖真实资源错误;当前 except RuntimeError 偏宽。
  3. 输出不再 bit-exact:PR body 说明后端切换在 bf16 下非 bit-exact,kernel 级 max-abs-diff 为 2.4e-4(self)/ 7.8e-3(cross),50-step 迭代去噪会放大差异;同 seed 全视频对比 PSNR 28.54 dB、SSIM 0.9468,属于用户使用 --attention-backend 已能获得的自由度,但依赖严格数值复现的流水线需要知晓。
  4. 量化模型数值保护依赖新逻辑_default_quantized_attention_backend 仅在 is_modelopt_fp4 时生效,若未来新增其他量化格式(如 FP8)也需要类似保护,当前判断口径较窄,需要后续扩展时保持同步。
  5. 回归面transformer_loader.py 新增 context manager 包裹 FSDP 加载,若量化默认解析出错可能影响模型构造期;已有单测覆盖 Blackwell 与显式后端两条路径。

影响范围集中在 sglang.multimodal_gen 的 sm_100 平台:B200 上运行 Wan2.2-TI2V/A14B、LingBot-World、MiniMax-H3、Qwen-Image、FLUX 等扩散模型的用户会自动获得约 1.13-1.5x 的端到端或 kernel 级收益,其中长序列(LingBot 506K)收益显著。系统层面,默认值变更不引入新配置项,用户可通过显式 --attention-backend fa 恢复旧行为;ModelOpt NVFP4 量化模型被单独排除在 cuDNN 默认之外,保证数值稳定。对 CI 而言,H100/Hopper 和 sm_120 的 golden 输出不受影响(construction 上仅 cc 10.x 生效),但团队需要关注后续测试矩阵中 B200 job 的 perf 阈值是否要被新的更快基线重新标定。

sm_100 默认后端变更 NVFP4 数值保护 运行时 fallback 依赖 RuntimeError 非 bit-exact 输出 H100 回归误报已澄清

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论