执行摘要
- 一句话:为 gfx950 添加 FP8 ASM prefill 加速 MLA 后端
- 推荐动作:推荐精读该 PR,尤其关注:自动检测函数
_fp8_mla_prefill_supported 的容错设计、持久调度元数据预分配策略、以及 Workspace Manager 的集成方式。Review 讨论中 host-device sync 的消除技巧值得学习。
功能与动机
在 gfx950 上为 DeepSeek-V3 等 MLA 模型降低 prefill 延迟并提高吞吐量。当 AITER 提供专用 FP8 ASM 内核时,代替 flash_attn_varlen_func 执行预填充,减少计算开销。
实现拆解
- 自动检测 (
_fp8_mla_prefill_supported):在模块加载时检测是否 gfx950 且 AITER 导出 mla_prefill_ps_asm_fwd 和 mla_reduce_v1,结果缓存避免重复。
- 预分配持久调度缓冲区 (
_init_fp8_prefill_ps_buffers):在 AiterMLAMetadataBuilder.__init__ 中根据 max_model_len 和 max_num_batched_tokens 计算最大预填充序列长度,调用 get_ps_metadata_info_v1 分配固定大小缓冲区。
- 元数据构建 (
_build_fp8_prefill_ps_metadata):在 build 方法中如果启用且存在 prefill 请求,调用该函数填充 PS 元数据,通过 common_attn_metadata 的 CPU 版本避免 host-device 同步。
- 前向分发 (
forward_mha):在 prefill 分支判断是否启用且非分块上下文(chunked-prefill 回退到 flash),调用 _mla_fp8_prefill_attn 并传递现有 output 避免额外拷贝。
- 内核调用 (
_mla_fp8_prefill_attn):先将 Q/K/V 从 bf16 显式转换为 fp8,然后调用 mla_prefill_ps_asm_fwd 生成部分和,再调用 mla_reduce_v1 归约到最终输出。临时 scratch 由 workspace manager 管理。
- 验证:在 MI355X TP=4 上通过
vllm bench serve 和 lm_eval gsm8k 验证性能和准确率。
关键文件:
vllm/v1/attention/backends/mla/rocm_aiter_mla.py(模块 MLA后端;类别 source;类型 core-logic;符号 _fp8_mla_prefill_supported, _init_fp8_prefill_ps_buffers, _build_fp8_prefill_ps_metadata, _mla_fp8_prefill_attn): 核心实现文件,新增 369 行实现 FP8 ASM prefill 的自动检测、缓冲区预分配、元数据构建和前向分发。
关键符号:_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
核心实现文件,新增 369 行实现 FP8 ASM prefill 的自动检测、缓冲区预分配、元数据构建和前向分发。
# -*- coding: utf-8 -*-
# SPDX-License-Identifier: Apache-2.0
import functools
from vllm.logger import init_logger
logger = 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,
)
评论区精华
风险与影响
- 风险:
- 平台限制:仅在 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
关联脉络
参与讨论