Prhub

#20319 [AMD] Support fp8 MLA for diffusion model

原始 PR 作者 yichiche 合并时间 2026-05-08 15:56 文件变更 1 提交数 7 评论 11 代码增减 +302 / -1

执行摘要

AMD 扩散模型 FP8 MLA 注意力优化,加速 MI355X 推理

PR描述中明确提出:"Replace the FP8 per-tensor flash attention path with the FP8 MLA prefill ASM kernel, which delivers significantly better performance on MI355X." 此外,依赖外部库ROCm/aiter PR#2481的升级。

值得精读,尤其是 _build_mla_prefill_metadata_can_use_mla_prefill 的实现,它们展示了如何将一个为DeepSeek设计的ASM内核(MLA)适配到扩散模型的DiT注意力中。关注 @torch.compiler.disable 的后续拆分计划。该PR为AMD平台DIffusion模型树立了FP8注意力优化的范例。

讨论亮点

review中的核心讨论围绕 @torch.compiler.disable 装饰器的使用:

  • avjves 指出将 @torch.compiler.disable 施加在整个 forward 方法上,会禁用BF16路径的编译优化,建议只门控MLA/FP8分支。
  • yichiche 回应这是保守选择,因为该路径混合了自定义aiter扩展、运行时形状分支和量化逻辑,当前解决方案简单可靠;承诺作为Future Work改进。
  • 最终决定保留全局disable,以待后续重构。

实现拆解

实现分为以下几步:

  1. 环境变量门控与GPU架构检查:新增 SGLANG_AITER_FP8_ATTN 环境变量控制FP8注意力开关,导入 _use_aiter_gfx95 函数检测是否为MI350/MI355(gfx950);定义模块级常量 _MLA_PREFILL_V_HEAD_DIM=128_MLA_PREFILL_HEAD_TILE=8 记录内核约束。

  2. 形状安全检查:实现 _can_use_mla_prefill(v_head_dim, num_heads) 函数,仅在GPU架构为gfx950、V头维度为128且头数能被8整除时返回True;否则走回退路径。

  3. 持久调度元数据构建:实现 _build_mla_prefill_metadata 函数,根据批次大小、序列长度、头数等参数构造MLA预填充内核所需的indptr、indices和work/reduce分区元数据;从SGLang SRTAiter后端模式改编而来。

  4. MLA内核调用与零填充:在 AITerImpl.forward 方法中新增分支:当启用FP8且形状允许时,先将Q/K从原生头维度(如128)零填充到192维,调用 mla_prefill_ps_asm_fwdmla_reduce_v1 内核,然后移除填充;否则使用原始的 aiter.flash_attn_func(BF16)。

  5. 测试与配置:未包含单元测试;配置上通过环境变量 SGLANG_AITER_FP8_ATTN=1 启用,需结合aiter升级版本使用。

文件 模块 状态 重要度
python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py 扩散模型注意力 modified 8.62

关键符号

_can_use_mla_prefill _build_mla_prefill_metadata _mla_prefill_ps_attention

关键源码片段

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

核心变更文件,新增 FP8 MLA 预填充 ASM 内核支持,通过环境变量门控,包含形状检查、元数据构建和内核调用逻辑。

# SPDX-License-Identifier: Apache-2.0import logging
import os
from typing import Optionalimport aiter
import torchfrom sglang.srt.models.deepseek_common.utils import _use_aiter_gfx95logger = logging.getLogger(__name__)_use_fp8_attn = os.environ.get("SGLANG_AITER_FP8_ATTN", "0") == "1"
_fp8_dtype = torch.float8_e4m3fn# ── MLA prefill ASM kernel constraints ──────────────────────────────
# 硬编码的内核约束(仅 gfx950 可用)
_MLA_PREFILL_QK_HEAD_DIM = 192
_MLA_PREFILL_V_HEAD_DIM = 128
_MLA_PREFILL_HEAD_TILE = 8if _use_fp8_attn:
    logger.info("DiT FP8 attention enabled via SGLANG_AITER_FP8_ATTN=1")def _can_use_mla_prefill(v_head_dim: int, num_heads: int) -> bool:
    """Check if the MLA prefill ASM kernel supports the given shape and GPU."""
    return (
        _use_aiter_gfx95 # AMD gfx950 架构检测
        and v_head_dim == _MLA_PREFILL_V_HEAD_DIM # V 头维度必须为 128
        and num_heads % _MLA_PREFILL_HEAD_TILE == 0 # 头数必须能被 8 整除
    )def _build_mla_prefill_metadata(
    batch_size: int,
    seq_lens: torch.Tensor,
    num_heads: int,
    num_kv_heads: int,
    is_causal: bool,
    block_size: int = 1,
    tile_q: int = 256,
    tile_kv: int = 128,
    kv_seq_lens: Optional[torch.Tensor] = None,
) -> dict:
    """Build persistent-scheduling metadata required by mla_prefill_ps_asm_fwd."""
    if kv_seq_lens is None:
        kv_seq_lens = seq_lens
​
    device = "cuda"
    gqa_ratio = num_heads // num_kv_heads
​
    qo_indptr = torch.zeros(batch_size + 1, dtype=torch.int32, device=device)
    kv_indptr = torch.zeros(batch_size + 1, dtype=torch.int32, device=device)
​
    qo_indptr[1 : batch_size + 1] = torch.cumsum(seq_lens, dim=0)
    actual_blocks = (kv_seq_lens + block_size - 1) // block_size
    kv_indptr[1 : batch_size + 1] = torch.cumsum(actual_blocks, dim=0)
    num_blocks = int(kv_indptr[-1])
​
    kv_indices = torch.arange(num_blocks, dtype=torch.int32, device=device)
    max_qlen = seq_lens.max()
    # ( 后续 work 分区与 reduce map 计算省略 )
    return {
        "qo_indptr": qo_indptr,
        "kv_indptr": kv_indptr,
        "kv_indices": kv_indices,
        # ... 其他元数据 ...
    }

评论区精华

torch.compiler.disable 覆盖整个 forward 的合理性 性能

avjves 指出 @torch.compiler.disable 禁用了 BF16 路径的编译优化;作者 yichiche 回应这是保守选择,当前路径混合了自定义扩展和运行时形状分支。

结论:暂时保留全局 disable,作为 future work 后续拆分。 · 已解决

风险与影响

具体风险:

  1. 架构依赖:ASM内核仅适用于MI355X(gfx950),其他AMD GPU(如MI300X)即使设置了环境变量也会回退,但若 _use_aiter_gfx95 误检可能崩溃。
  2. 形状约束num_heads % 8 != 0v_head_dim != 128 的组合可能导致内核OOB读取或错误结果;当前检查已覆盖,但缺少测试验证。
  3. 外部依赖:依赖ROCm/aiter PR#2481的特定版本,若未升级则导入失败或行为不一致。
  4. 编译影响:全局 @torch.compiler.disable 阻止了BF16路径的图优化,可能影响其他模型性能。

影响范围:

  • 用户:仅影响使用AMD GPU(特别是MI355X)运行扩散模型(如Wan2.2)的用户。启用FP8后单卡81帧720p视频生成总时间减少约19%。
  • 系统:新代码仅在设置 SGLANG_AITER_FP8_ATTN=1_use_aiter_gfx95 为True时激活,默认不开,对其他模型无影响。
  • 团队:AMD和扩散模型团队需要维护新增的MLA内核调用和元数据构建逻辑,以及形状约束表。
依赖外部 aiter 版本 缺少单元测试 仅 MI355X 支持 全局 torch.compiler.disable

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论