Prhub

#42509 [ROCm][MLA] FP8 ASM prefill for AITER dense MLA backend on gfx950

原始 PR 作者 maeehart 合并时间 2026-05-15 23:56 文件变更 1 提交数 5 评论 35 代码增减 +369 / -0

执行摘要

为 gfx950 添加 FP8 ASM prefill 加速 MLA 后端

在 gfx950 上为 DeepSeek-V3 等 MLA 模型降低 prefill 延迟并提高吞吐量。当 AITER 提供专用 FP8 ASM 内核时,代替 flash_attn_varlen_func 执行预填充,减少计算开销。

推荐精读该 PR,尤其关注:自动检测函数 _fp8_mla_prefill_supported 的容错设计、持久调度元数据预分配策略、以及 Workspace Manager 的集成方式。Review 讨论中 host-device sync 的消除技巧值得学习。

讨论亮点
  • 性能同步问题:gemini-code-assist 指出 _build_fp8_prefill_ps_metadata 中的 .to("cpu")._mla_fp8_prefill_attn 中的 .item() 引入不必要的 host-device 同步,阻碍 CUDA Graph 捕获。作者改为使用 common_attn_metadata.query_start_loc_cpu 切片,并将 .item() 移至元数据构建阶段缓存。
  • FP8 显式转换必要性:gemini 认为内核内部可处理转换,作者澄清 mla_prefill_ps_asm_fwd 期望 fp8 输入,而 Q/K/V 从 kv_b_proj 输出为 bf16,必须显式转换,已在注释中更正。
  • 输出 Buffer 优化:gemini 建议 _mla_fp8_prefill_attn 直接使用传入的 output 避免临时分配和 copy_。作者采纳,签名改为 _mla_fp8_prefill_attn(..., out: torch.Tensor)
  • 初始化简化:tjtanaa 建议移除冗余除法,改由注释保留语义。作者已将 gqa_ratio 等运算替换为常量和注释。
  • Workspace Manager:tjtanaa 建议遵循 PR #41002 模式管理临时 scratch。作者为 logitsattn_lse 等 scratch 改用 current_workspace_manager().get_simultaneous

实现拆解

  1. 自动检测 (_fp8_mla_prefill_supported):在模块加载时检测是否 gfx950 且 AITER 导出 mla_prefill_ps_asm_fwdmla_reduce_v1,结果缓存避免重复。
  2. 预分配持久调度缓冲区 (_init_fp8_prefill_ps_buffers):在 AiterMLAMetadataBuilder.__init__ 中根据 max_model_lenmax_num_batched_tokens 计算最大预填充序列长度,调用 get_ps_metadata_info_v1 分配固定大小缓冲区。
  3. 元数据构建 (_build_fp8_prefill_ps_metadata):在 build 方法中如果启用且存在 prefill 请求,调用该函数填充 PS 元数据,通过 common_attn_metadata 的 CPU 版本避免 host-device 同步。
  4. 前向分发 (forward_mha):在 prefill 分支判断是否启用且非分块上下文(chunked-prefill 回退到 flash),调用 _mla_fp8_prefill_attn 并传递现有 output 避免额外拷贝。
  5. 内核调用 (_mla_fp8_prefill_attn):先将 Q/K/V 从 bf16 显式转换为 fp8,然后调用 mla_prefill_ps_asm_fwd 生成部分和,再调用 mla_reduce_v1 归约到最终输出。临时 scratch 由 workspace manager 管理。
  6. 验证:在 MI355X TP=4 上通过 vllm bench servelm_eval gsm8k 验证性能和准确率。
文件 模块 状态 重要度
vllm/v1/attention/backends/mla/rocm_aiter_mla.py MLA 后端 modified 8.65

关键符号

_fp8_mla_prefill_supported _init_fp8_prefill_ps_buffers _build_fp8_prefill_ps_metadata _mla_fp8_prefill_attn forward_mha

关键源码片段

vllm/v1/attention/backends/mla/rocm_aiter_mla.py core-logic

核心实现文件,新增 369 行实现 FP8 ASM prefill 的自动检测、缓冲区预分配、元数据构建和前向分发。

# -*- coding: utf-8 -*-
# SPDX-License-Identifier: Apache-2.0import functools
from vllm.logger import init_loggerlogger = init_logger(__name__)
​
​
@functools.lru_cache(maxsize=1)
def _fp8_mla_prefill_supported() -> bool:
    """Auto-detect FP8 MLA prefill support on gfx950.    Checks both hardware and AITER kernel availability.
    Result is cached to avoid repeated import overhead.
    """
    try:
        from vllm.platforms.rocm import on_gfx950
    except Exception:
        return False
    if not on_gfx950():
        return False
    try:
        # FP8 ASM kernels for prefill, packaged in AITER
        from aiter import mla_prefill_ps_asm_fwd, mla_reduce_v1 # noqa: F401
    except Exception:
        return False
    return True
​
​
# In AiterMLAImpl.forward_mha, prefill dispatch:
if attn_metadata.prefill is not None and self.fp8_prefill_enabled:
    # Use FP8 ASM kernel, but skip chunked-prefill (fallback to flash)
    if not attn_metadata.prefill.is_chunked:
        self._mla_fp8_prefill_attn(q, k, v, attn_metadata, output)
    else:
        self._flash_attn_varlen_func(q, k, v, output)
else:
    # Standard decode or fallback
    ...
​
​
# In _mla_fp8_prefill_attn, final two-stage execution:
def _mla_fp8_prefill_attn(
    self,
    q: torch.Tensor, # bf16, shape [total_q, nhead, v_head_dim]
    k: torch.Tensor,
    v: torch.Tensor,
    attn_metadata: AiterMLAMetadata,
    out: torch.Tensor, # pre-allocated output buffer
) -> None:
    # Step 1: explicit cast from bf16 to fp8 (kernel expects fp8 inputs)
    # q_scale/k_scale/v_scale = 1.0 disables internal scaling
    q_fp8 = q.to(torch.float8_e4m3fn)
    k_fp8 = k.to(torch.float8_e4m3fn)
    v_fp8 = v.to(torch.float8_e4m3fn)
​
    total_q, nhead, v_head_dim = q.shape
    out_3d = out.view(total_q, nhead, v_head_dim)
​
    # Step 2: persistent-scheduling forward
    self._mla_prefill_ps_asm_fwd(
        q_fp8, k_fp8, v_fp8,
        attn_metadata.fp8_prefill_qo_indptr,
        attn_metadata.fp8_prefill_kv_indptr,
        attn_metadata.fp8_prefill_kv_indices,
        attn_metadata.fp8_prefill_work_indptr,
        attn_metadata.fp8_prefill_work_info_set,
        out_3d, # writes partital results directly
        q_scale=1.0,
        k_scale=1.0,
        v_scale=1.0,
    )
​
    # Step 3: reduction
    self._mla_reduce_v1(
        out_3d, # same buffer, reduced in-place
        attn_metadata.fp8_prefill_reduce_indptr,
        attn_metadata.fp8_prefill_reduce_final_map,
        attn_metadata.fp8_prefill_reduce_partial_map,
        attn_metadata.fp8_prefill_num_partial_tiles,
    )

评论区精华

避免 host-device 同步在 metadata 构建中 性能

gemini-code-assist 指出 `.to("cpu")` 引入同步,建议使用 `common_attn_metadata.query_start_loc_cpu`

结论:作者将 `_build_fp8_prefill_ps_metadata` 签名改为接受 `common_attn_metadata`,切片 CPU tensor 避免同步。 · 已解决

FP8 显式转换必要性 正确性

gemini-code-assist 认为内核内部会转换,建议移除显式 cast 以省开销

结论:作者澄清内核期望 FP8 输入,Q/K/V 来自 bf16 必须显式转换,更新注释并保留 cast。 · 已解决

使用 workspace manager 管理 scratch 内存 性能

tjtanaa 建议遵循 PR #41002 模式管理临时 scratch,减少分配开销

结论:作者为 `logits`, `attn_lse`, `final_lse` 改用 `current_workspace_manager().get_simultaneous`。 · 已解决

风险与影响

  • 平台限制:仅在 gfx950 + AITER 提供 FP8 ASM 内核时生效,其他平台自动回退,但若检测条件过于宽松可能在其他设备上引入错误(当前通过 on_gfx950 严格限制)。
  • 精度影响:FP8 预填充使用 one_scale=1.0 固定缩放,未按最大绝对值校准,可能在某些输入下精度受损。目前无环境变量强制禁用,用户需等待后续迭代(如 fxmarty-amd 建议的警告/开关)。
  • 缺少测试覆盖:无自动化单元测试,仅依赖手动验证和 benchmark,回归风险较大。
  • 同步点消除不完全:虽然已移除 .to("cpu").item(),但仍需确认 get_ps_metadata_info_v1 等调用无隐藏同步。
  • 用户影响:MI355X 上使用 --kv-cache-dtype fp8 的 DeepSeek-V3 用户自动获得 TTFT 降低 14.8%,吞吐量提升 2.3%,无需任何配置。BF16 KV cache 用户同样受益(预填充内核独立于 KV 存储格式)。
  • 系统影响:通过卸载至专用 ASM 内核,减少 CC 延迟和显存分配,降低 GPU 压力。
  • 团队影响:为 ROCm MLA 后端建立硬件特定优化模式,未来可扩展至其他内核(如 decode)。
缺少测试覆盖 精度依赖默认缩放 仅 gfx950

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论