Prhub

#34107 [diffusion] fix: guard SageAttention SM90 bindings

原始 PR 作者 RunFMe 合并时间 2026-08-08 21:32 文件变更 3 提交数 3 评论 2 代码增减 +53 / -4

执行摘要

SageAttention 补 SM90 守卫,缺补丁时回退 FA

PR body 指出 sageattention==2.2.0 lacks the upstream SM90 binding fix from thu-ml/SageAttention#307。On Hopper it can return uninitialized attention output without launching the attention kernel。选择绑定形状检测而不是包版本,是因为 fixed upstream source still reports version 2.2.0,版本号无法区分打补丁和未打补丁的 PyPI 发布。

值得快速阅读:这是一个典型的「可选依赖 fail-closed 守卫」案例,亮点在于用绑定形状(符号是否存在)而非包版本做检测,因为上游修复后版本号仍为 2.2.0。对于需要在源码里探测「同版本不同实现」的场景有借鉴价值;同时可关注其 CI Extra 失败重跑后的结果,以及未来是否补上正向路径测试。

讨论亮点

本 PR 没有 review 级代码评论,仅有合入流程相关的两条评论:RunFMe 在关联 Issue 下 @mickqian "could you take a look pls?",请求维护者确认方案;mickqian 随后执行 /rerun-failed-ci,对 Extra CI 失败(Run #31255494471)进行重跑,PR 最终由 mickqian 合并。设计取舍(为何按符号而非版本检测)主要在 PR body 中说明,没有引发反对意见。

实现拆解

  1. 变更入口:在 python/sglang/multimodal_gen/runtime/platforms/cuda.py 中重写 _SageAttentionBackendResolver.resolve。原实现将「导入 sageattention 包」和「导入 SageAttentionBackend」混在一个 try 中;新实现分三步:先只导入 sageattention 包,若导入失败记录提示并返回 AttentionBackendEnum.FA;之后若 platform.is_hopper() 为真,尝试导入 sageattention.sm90_compile 中的 qk_int8_sv_f8_accum_f32_fuse_v_scale_attn_inst_buf_fake_impl(上游 SM90 修复保留的假实现符号),导入失败即判定为未打补丁的 2.2.0,记录 warning 后返回 AttentionBackendEnum.FA;最后才导入 SageAttentionBackend 并返回其类字符串,导入失败同样回退 FA。

  2. 关键取舍:安装提示从 pip install sageattention==2.2.0 改为 pip install git+https://github.com/thu-ml/SageAttention.git@d9704247a5139ab4c03bf7fc6b35cc0e2cbb5ea4 --no-build-isolation,以固定提交保证拿到 SM90 绑定修复。

  3. 测试配套:在 python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py 中为 FakeCudaPlatform 增加 is_hopper 类方法和 is_hopper_device 状态,并在 setUp 中重置;新增 test_hopper_sage_attention_without_sm90_fix_falls_back,用 types.ModuleType 构造不含修复符号的假 sageattention 模块,通过 patch.dict(sys.modules, ...) 注入后断言 _SageAttentionBackendResolver.resolve 返回 AttentionBackendEnum.FA

  4. 文档配套:docs/docs/sglang-diffusion/attention_backends.mdx 中更新 sage_attn 表格行,注明 Hopper 上 PyPI 2.2.0 不受支持,并给出固定提交的安装命令;fallback 说明也补充「Hopper 上缺少 SM90 绑定修复时同样回退 FlashAttention」。

文件 模块 状态 重要度
python/sglang/multimodal_gen/runtime/platforms/cuda.py 平台层 modified 6.32
python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py 注意力后端 modified 5.68
docs/docs/sglang-diffusion/attention_backends.mdx 功能文档 modified 2.49

关键符号

_SageAttentionBackendResolver.resolve is_hopper test_hopper_sage_attention_without_sm90_fix_falls_back

关键源码片段

python/sglang/multimodal_gen/runtime/platforms/cuda.py dependency-wiring

核心变更文件:在 _SageAttentionBackendResolver.resolve 中拆分流式导入与后端导入,加入 Hopper SM90 绑定守卫,缺符号时 fail-closed 到 FlashAttention,并更新安装提示。

class _SageAttentionBackendResolver(_CudaAttentionBackendResolver):
    backend = AttentionBackendEnum.SAGE_ATTN
​
    @classmethod
    def resolve(cls, platform) -> str | AttentionBackendEnum:
        # 第一步:确认 SageAttention 包本身可导入(PyPI 上发布的是 2.2.0)
        try:
            from sageattention import sageattn # noqa: F401
        except ImportError as e:
            logger.info(e)
            logger.info(
                "Sage Attention backend is not installed (To install it, run "
                "`pip install git+https://github.com/thu-ml/SageAttention.git"
                "@d9704247a5139ab4c03bf7fc6b35cc0e2cbb5ea4 --no-build-isolation`). "
                "Falling back to Flash Attention."
            )
            return AttentionBackendEnum.FA
​
        # 第二步:Hopper 上额外检查 SM90 绑定修复是否存在。
        # 上游修复(thu-ml/SageAttention PR #307)保留了名为 fake impl 的符号,
        # 而 PyPI 2.2.0 没有该符号;版本号仍是 2.2.0,所以只能用符号来区分。
        if platform.is_hopper():
            try:
                from sageattention.sm90_compile import ( # noqa: F401
                    qk_int8_sv_f8_accum_f32_fuse_v_scale_attn_inst_buf_fake_impl,
                )
            except ImportError:
                # 没有补丁时,Hopper 上可能不启动内核就返回未初始化结果,
                # fail closed 到 FlashAttention 比静默出错更安全。
                logger.warning(
                    "Installed Sage Attention is missing the SM90 binding fix. "
                    "Falling back to Flash Attention. Reinstall with "
                    "`pip install --force-reinstall git+https://github.com/thu-ml/"
                    "SageAttention.git@d9704247a5139ab4c03bf7fc6b35cc0e2cbb5ea4 "
                    "--no-build-isolation`."
                )
                return AttentionBackendEnum.FA
​
        # 第三步:前两步通过后再导入 SGLang 侧后端实现。
        try:
            from sglang.multimodal_gen.runtime.layers.attention.backends.sage_attn import ( # noqa: F401
                SageAttentionBackend,
            )
            return (
                "sglang.multimodal_gen.runtime.layers.attention.backends.sage_attn."
                "SageAttentionBackend"
            )
        except ImportError as e:
            logger.info(e)
            logger.info(
                "Sage Attention backend failed to import. Falling back to Flash Attention."
            )
            return AttentionBackendEnum.FA
python/sglang/multimodal_gen/test/unit/test_cuda_attention_backend.py test-coverage

新增 Hopper fail-closed 路径的单元测试:为 FakeCudaPlatform 补充 is_hopper,并用 patch.dict(sys.modules) 伪造无修复符号的 sageattention 模块,验证解析结果回退 FA。

class FakeCudaPlatform(CudaPlatformBase):
    is_sm120_device = False
    is_blackwell_device = False
    is_hopper_device = False
    supports_flash_attention = True
​
    @classmethod
    def is_sm120(cls):
        return cls.is_sm120_device
​
    @classmethod
    def is_blackwell(cls):
        return cls.is_blackwell_device
​
    # 为覆盖 Hopper 守卫路径而新增的设备能力判断。
    @classmethod
    def is_hopper(cls):
        return cls.is_hopper_device
​
    @classmethod
    def has_device_capability(
        cls,
        capability: tuple[int, int] | int,
        device_id: int = 0,
    ) -> bool:
        return cls.supports_flash_attention
​
​
def test_hopper_sage_attention_without_sm90_fix_falls_back(self):
    FakeCudaPlatform.is_hopper_device = True
    # 用内存模块伪造没有修复符号的 sageattention,避免真实安装干扰测试。
    sageattention = types.ModuleType("sageattention")
    sageattention.__path__ = []
    sageattention.sageattn = object()
    sm90_compile = types.ModuleType("sageattention.sm90_compile")
​
    with patch.dict(
        sys.modules,
        {"sageattention": sageattention, "sageattention.sm90_compile": sm90_compile},
    ):
        # 缺少 qk_int8_sv_f8_accum_f32_fuse_v_scale_attn_inst_buf_fake_impl 时,
        # Hopper 上必须回退到 FlashAttention。
        self.assertEqual(
            _SageAttentionBackendResolver.resolve(FakeCudaPlatform),
            AttentionBackendEnum.FA,
        )

评论区精华

SageAttention 守卫合入与失败 CI 重跑 question

RunFMe 在关联 Issue 下 @mickqian "could you take a look pls?",请求维护者确认安装守护方案;mickqian 随后回复 "/rerun-failed-ci" 并重跑 Extra CI(Run #31255494471)。

结论:维护者确认后合入;没有设计层面的反对意见,但 Extra CI 曾有一次失败运行,合并前已重跑。 · 已解决

风险与影响

  1. 依赖固定 commit hash:安装引导指向 d9704247...,若该提交在上游仓库被 rebase 或 GC,安装命令会失效,后续需更新 hash。
  2. 基于上游内部符号检测:守卫依赖 sageattention.sm90_compile.qk_int8_sv_f8_accum_f32_fuse_v_scale_attn_inst_buf_fake_impl 的存在性,符号名是上游内部实现细节,改名或重构会导致守卫误判(过度回退 FA,虽安全但损失性能)。
  3. 覆盖面仅限于 Hopper:platform.is_hopper() 分支只保护 SM90,Blackwell(SM120)等其他架构若存在类似问题不在本次修复范围。
  4. 测试只覆盖 fail-closed 路径:缺少「打补丁后正常选择 SageAttentionBackend」的正向测试,后续演进时回归风险需由其他测试兜底。
  5. Extra CI 曾有一次失败运行(#31255494471),合并前已重跑,但失败根因未在评论中说明,需关注是否与本改动相关。

影响范围集中在扩散模块(multimodal_gen)的可选注意力后端:Hopper 用户从「可能拿到未初始化注意力输出」变为「明确警告并回退 FlashAttention」,正确性显著提升,代价是未打补丁的 PyPI 2.2.0 在 Hopper 上不再可用 SageAttention。对 Blackwell、Ampere 等非 Hopper 用户无行为变化。仓库层面同步更新了安装文档和单元测试,降低了后续维护者误用 PyPI 2.2.0 的概率。

固定 commit hash 安装指引 基于上游内部符号检测 守卫仅覆盖 Hopper 正向路径缺测试 Extra CI 曾失败

关联 Issue

#307 fix same name issue in sm09_compile.py

完整报告

参与讨论