执行摘要
- 一句话:AMD扩散模型FP8 MLA注意力优化,加速MI355X推理
- 推荐动作:值得精读,尤其是
_build_mla_prefill_metadata 和 _can_use_mla_prefill 的实现,它们展示了如何将一个为DeepSeek设计的ASM内核(MLA)适配到扩散模型的DiT注意力中。关注 @torch.compiler.disable 的后续拆分计划。该PR为AMD平台DIffusion模型树立了FP8注意力优化的范例。
功能与动机
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的升级。
实现拆解
实现分为以下几步:
-
环境变量门控与GPU架构检查:新增 SGLANG_AITER_FP8_ATTN 环境变量控制FP8注意力开关,导入 _use_aiter_gfx95 函数检测是否为MI350/MI355(gfx950);定义模块级常量 _MLA_PREFILL_V_HEAD_DIM=128 和 _MLA_PREFILL_HEAD_TILE=8 记录内核约束。
-
形状安全检查:实现 _can_use_mla_prefill(v_head_dim, num_heads) 函数,仅在GPU架构为gfx950、V头维度为128且头数能被8整除时返回True;否则走回退路径。
-
持久调度元数据构建:实现 _build_mla_prefill_metadata 函数,根据批次大小、序列长度、头数等参数构造MLA预填充内核所需的indptr、indices和work/reduce分区元数据;从SGLang SRTAiter后端模式改编而来。
-
MLA内核调用与零填充:在 AITerImpl.forward 方法中新增分支:当启用FP8且形状允许时,先将Q/K从原生头维度(如128)零填充到192维,调用 mla_prefill_ps_asm_fwd 和 mla_reduce_v1 内核,然后移除填充;否则使用原始的 aiter.flash_attn_func(BF16)。
-
测试与配置:未包含单元测试;配置上通过环境变量 SGLANG_AITER_FP8_ATTN=1 启用,需结合aiter升级版本使用。
关键文件:
python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py(模块 扩散模型注意力;类别 source;类型 core-logic;符号 _can_use_mla_prefill, _build_mla_prefill_metadata, _mla_prefill_ps_attention): 核心变更文件,新增FP8 MLA预填充ASM内核支持,通过环境变量门控,包含形状检查、元数据构建和内核调用逻辑。
关键符号:_can_use_mla_prefill, _build_mla_prefill_metadata, _mla_prefill_ps_attention
关键源码片段
python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py
核心变更文件,新增FP8 MLA预填充ASM内核支持,通过环境变量门控,包含形状检查、元数据构建和内核调用逻辑。
# SPDX-License-Identifier: Apache-2.0
import logging
import os
from typing import Optional
import aiter
import torch
from sglang.srt.models.deepseek_common.utils import _use_aiter_gfx95
logger = 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 = 8
if _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,
# ... 其他元数据 ...
}
评论区精华
review中的核心讨论围绕 @torch.compiler.disable 装饰器的使用:
风险与影响
- 架构依赖:ASM内核仅适用于MI355X(gfx950),其他AMD GPU(如MI300X)即使设置了环境变量也会回退,但若
_use_aiter_gfx95 误检可能崩溃。
- 形状约束:
num_heads % 8 != 0 或 v_head_dim != 128 的组合可能导致内核OOB读取或错误结果;当前检查已覆盖,但缺少测试验证。
- 外部依赖:依赖ROCm/aiter PR#2481的特定版本,若未升级则导入失败或行为不一致。
- 编译影响:全局
@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
关联脉络
- PR #23955 [AMD] Add AMD FP8 MLA attention test for Wan2.2-T2V-A14B: 此后续 PR 添加了针对本 PR 核心功能(FP8 MLA 注意力)的单元测试,覆盖 Wan2.2 模型。
参与讨论