[AMD] Replace fp8 mla with fp8 mha kernel for diffusion model aiter backend (#23927)
This commit is contained in:
@@ -4,7 +4,6 @@
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
import aiter
|
||||
import torch
|
||||
@@ -23,39 +22,66 @@ logger = logging.getLogger(__name__)
|
||||
_use_fp8_attn = os.environ.get("SGLANG_DIFFUSION_AITER_FP8_ATTN", "0") == "1"
|
||||
_fp8_dtype = torch.float8_e4m3fn
|
||||
|
||||
# ── MLA prefill ASM kernel constraints ──────────────────────────────
|
||||
# The only available FP8 prefill kernel is the pre-compiled ASM binary
|
||||
# mla_pfl_qh192_vh128_m32x8_n128x1_causal{0,1}.co, originally built for
|
||||
# DeepSeek-style MLA. Four hard constraints:
|
||||
#
|
||||
# 1. GPU arch must be gfx950 (MI350/MI355). The ASM binary is compiled
|
||||
# exclusively for gfx950; it will crash or fail to load on other archs.
|
||||
#
|
||||
# 2. qk_head_dim baked at 192. Models with smaller head dims (e.g.
|
||||
# Wan's 128) are handled by zero-padding Q/K — extra dims contribute
|
||||
# 0 to dot products, preserving correctness.
|
||||
#
|
||||
# 3. v_head_dim baked at 128. Models with V head dim != 128 cannot use
|
||||
# this kernel.
|
||||
#
|
||||
# 4. Kernel tiles over heads in groups of 8 ("m32x8" = 32 tokens x 8
|
||||
# heads per tile). num_heads not divisible by 8 causes OOB reads.
|
||||
# E.g. Ulysses SP degree=4 with 40 heads -> 10 heads/rank -> crash.
|
||||
_MLA_PREFILL_QK_HEAD_DIM = 192
|
||||
_MLA_PREFILL_V_HEAD_DIM = 128
|
||||
_MLA_PREFILL_HEAD_TILE = 8
|
||||
# fmha_fwd_hd128_fp8_gfx950 ASM kernel. Support full MHA with q/k/v head_dim == 128 -- e.g., Wan 2.2 self- and cross-attention.
|
||||
_FMHA_FP8_HEAD_DIM = 128
|
||||
|
||||
|
||||
if _use_fp8_attn:
|
||||
logger.info("DiT FP8 attention enabled via SGLANG_DIFFUSION_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
|
||||
and v_head_dim == _MLA_PREFILL_V_HEAD_DIM
|
||||
and num_heads % _MLA_PREFILL_HEAD_TILE == 0
|
||||
def _can_use_fmha_fp8_prefill(
|
||||
q_head_dim: int,
|
||||
k_head_dim: int,
|
||||
v_head_dim: int,
|
||||
num_heads: int,
|
||||
num_kv_heads: int,
|
||||
) -> bool:
|
||||
"""True if MHA q/k/v head_dim==128 on a gfx950-class arch."""
|
||||
if not _use_aiter_gfx95:
|
||||
return False
|
||||
if num_kv_heads != num_heads:
|
||||
return False
|
||||
return q_head_dim == k_head_dim == v_head_dim == _FMHA_FP8_HEAD_DIM
|
||||
|
||||
|
||||
def _fmha_fp8_prefill_attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float,
|
||||
is_causal: bool,
|
||||
q_scale: torch.Tensor,
|
||||
k_scale: torch.Tensor,
|
||||
v_scale: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
FP8 FMHA prefill via aiter.flash_attn_fp8_pertensor_func.
|
||||
|
||||
Expects q, k, v as (batch, seqlen, nheads, 128) FP8, contiguous.
|
||||
"""
|
||||
|
||||
def _ensure_fp8_descale(scale: torch.Tensor) -> torch.Tensor:
|
||||
"""Per-tensor descale as shape (1,) float32 for flash_attn_fp8_pertensor_func."""
|
||||
return scale.to(dtype=torch.float32).reshape(1).contiguous()
|
||||
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
q_descale = _ensure_fp8_descale(q_scale)
|
||||
k_descale = _ensure_fp8_descale(k_scale)
|
||||
v_descale = _ensure_fp8_descale(v_scale)
|
||||
|
||||
return aiter.flash_attn_fp8_pertensor_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
q_descale,
|
||||
k_descale,
|
||||
v_descale,
|
||||
causal=is_causal,
|
||||
softmax_scale=softmax_scale,
|
||||
window_size=(-1, -1, 0),
|
||||
)
|
||||
|
||||
|
||||
@@ -82,225 +108,6 @@ class AITerBackend(AttentionBackend):
|
||||
raise NotImplementedError("AITer backend does not have a metadata builder.")
|
||||
|
||||
|
||||
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.
|
||||
|
||||
Args:
|
||||
batch_size: number of sequences in the batch.
|
||||
seq_lens: [batch_size] int tensor with per-sequence Q lengths (on CPU).
|
||||
num_heads: number of query heads.
|
||||
num_kv_heads: number of KV heads.
|
||||
is_causal: whether causal masking is used.
|
||||
block_size: KV page size (1 for non-paged token-level layout).
|
||||
tile_q: Q tile size used by the kernel.
|
||||
tile_kv: KV tile granularity.
|
||||
kv_seq_lens: [batch_size] int tensor with per-sequence KV lengths (on CPU).
|
||||
If None, defaults to seq_lens (self-attention).
|
||||
|
||||
Returns:
|
||||
dict with all metadata tensors needed by the kernel + reduce.
|
||||
"""
|
||||
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)
|
||||
kv_indptr = torch.zeros(batch_size + 1, dtype=torch.int32)
|
||||
|
||||
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)
|
||||
|
||||
max_qlen = seq_lens.max()
|
||||
|
||||
qhead_granularity = gqa_ratio
|
||||
qlen_granularity = tile_q // qhead_granularity
|
||||
kvlen_granularity = max(tile_kv, block_size)
|
||||
|
||||
(
|
||||
(work_meta_data_size, work_meta_data_type),
|
||||
(work_indptr_size, work_indptr_type),
|
||||
(work_info_size, work_info_type),
|
||||
(reduce_indptr_size, reduce_indptr_type),
|
||||
(reduce_final_map_size, reduce_final_map_type),
|
||||
(reduce_partial_map_size, reduce_partial_map_type),
|
||||
) = aiter.get_ps_metadata_info_v1(
|
||||
batch_size=batch_size,
|
||||
num_head_k=num_kv_heads,
|
||||
max_qlen=max_qlen,
|
||||
qlen_granularity=qlen_granularity,
|
||||
)
|
||||
|
||||
work_metadata_ptrs = torch.empty(
|
||||
work_meta_data_size, dtype=work_meta_data_type, device=device
|
||||
)
|
||||
work_indptr = torch.empty(work_indptr_size, dtype=work_indptr_type, device=device)
|
||||
work_info = torch.empty(work_info_size, dtype=work_info_type, device=device)
|
||||
reduce_indptr = torch.empty(
|
||||
reduce_indptr_size, dtype=reduce_indptr_type, device=device
|
||||
)
|
||||
reduce_final_map = torch.empty(
|
||||
reduce_final_map_size, dtype=reduce_final_map_type, device=device
|
||||
)
|
||||
reduce_partial_map = torch.empty(
|
||||
reduce_partial_map_size, dtype=reduce_partial_map_type, device=device
|
||||
)
|
||||
|
||||
aiter.get_ps_metadata_v1(
|
||||
qo_indptr.cpu(),
|
||||
kv_indptr.cpu(),
|
||||
seq_lens.cpu().int(),
|
||||
gqa_ratio,
|
||||
num_kv_heads,
|
||||
work_metadata_ptrs,
|
||||
work_indptr,
|
||||
work_info,
|
||||
reduce_indptr,
|
||||
reduce_final_map,
|
||||
reduce_partial_map,
|
||||
qhead_granularity=qhead_granularity,
|
||||
qlen_granularity=qlen_granularity,
|
||||
kvlen_granularity=kvlen_granularity,
|
||||
block_size=block_size,
|
||||
is_causal=is_causal,
|
||||
)
|
||||
|
||||
return {
|
||||
"qo_indptr": qo_indptr.to(device),
|
||||
"kv_indptr": kv_indptr.to(device),
|
||||
"kv_indices": kv_indices.to(device),
|
||||
"work_indptr": work_indptr,
|
||||
"work_info": work_info,
|
||||
"reduce_indptr": reduce_indptr,
|
||||
"reduce_final_map": reduce_final_map,
|
||||
"reduce_partial_map": reduce_partial_map,
|
||||
"max_seqlen_q": max_qlen,
|
||||
"tile_q": tile_q,
|
||||
}
|
||||
|
||||
|
||||
@torch.compiler.disable
|
||||
def _mla_prefill_ps_attention(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
softmax_scale: float,
|
||||
is_causal: bool,
|
||||
q_scale: Optional[torch.Tensor] = None,
|
||||
k_scale: Optional[torch.Tensor] = None,
|
||||
v_scale: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Run mla_prefill_ps_asm_fwd + mla_reduce_v1 on 4D batch tensors.
|
||||
|
||||
Reshapes [B, S, H, D] -> varlen [B*S, H, D] with trivial indptr,
|
||||
calls the kernel, then reshapes back.
|
||||
|
||||
Supports cross-attention where q has seq_len S_q and k/v have seq_len S_kv.
|
||||
|
||||
The ASM kernel has qk_head_dim=192 baked in at compile time (see module-level
|
||||
comments). If the model's head dim is smaller (e.g. 128), Q and K are
|
||||
zero-padded along the last dimension to 192 before calling the kernel.
|
||||
The padded zeros contribute nothing to the QK dot product, so attention
|
||||
scores are identical to the unpadded case.
|
||||
"""
|
||||
B, S_q, H, D_q = q.shape
|
||||
S_kv = k.shape[1]
|
||||
D_v = v.shape[-1]
|
||||
device = q.device
|
||||
num_kv_heads = k.shape[2]
|
||||
|
||||
# Zero-pad Q/K head dim to match the kernel's compiled qk_head_dim=192.
|
||||
# Padding with zeros preserves dot-product correctness.
|
||||
pad_qk = _MLA_PREFILL_QK_HEAD_DIM - D_q
|
||||
if pad_qk > 0:
|
||||
q = torch.nn.functional.pad(q, (0, pad_qk))
|
||||
k = torch.nn.functional.pad(k, (0, pad_qk))
|
||||
D_q_kernel = q.shape[-1]
|
||||
|
||||
q_varlen = q.reshape(B * S_q, H, D_q_kernel).contiguous()
|
||||
k_varlen = k.reshape(B * S_kv, num_kv_heads, D_q_kernel).contiguous()
|
||||
v_varlen = v.reshape(B * S_kv, num_kv_heads, D_v).contiguous()
|
||||
|
||||
q_seq_lens = torch.full((B,), S_q, dtype=torch.int32)
|
||||
kv_seq_lens = torch.full((B,), S_kv, dtype=torch.int32)
|
||||
|
||||
meta = _build_mla_prefill_metadata(
|
||||
batch_size=B,
|
||||
seq_lens=q_seq_lens,
|
||||
kv_seq_lens=kv_seq_lens,
|
||||
num_heads=H,
|
||||
num_kv_heads=num_kv_heads,
|
||||
is_causal=is_causal,
|
||||
block_size=1,
|
||||
)
|
||||
|
||||
total_s = B * S_q
|
||||
tile_q = meta["tile_q"]
|
||||
|
||||
output = torch.empty((total_s, H, D_v), dtype=torch.bfloat16, device=device)
|
||||
logits = torch.empty(
|
||||
(meta["reduce_partial_map"].size(0) * tile_q, H, D_v),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
attn_lse = torch.empty(
|
||||
(meta["reduce_partial_map"].size(0) * tile_q, H),
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
)
|
||||
final_lse = torch.empty((total_s, H), dtype=torch.float32, device=device)
|
||||
|
||||
aiter.mla_prefill_ps_asm_fwd(
|
||||
q_varlen,
|
||||
k_varlen,
|
||||
v_varlen,
|
||||
meta["qo_indptr"],
|
||||
meta["kv_indptr"],
|
||||
meta["kv_indices"],
|
||||
meta["work_indptr"],
|
||||
meta["work_info"],
|
||||
meta["max_seqlen_q"],
|
||||
softmax_scale,
|
||||
is_causal,
|
||||
logits,
|
||||
attn_lse,
|
||||
output,
|
||||
q_scale,
|
||||
k_scale,
|
||||
v_scale,
|
||||
)
|
||||
|
||||
aiter.mla_reduce_v1(
|
||||
logits,
|
||||
attn_lse,
|
||||
meta["reduce_indptr"],
|
||||
meta["reduce_final_map"],
|
||||
meta["reduce_partial_map"],
|
||||
tile_q,
|
||||
output,
|
||||
final_lse,
|
||||
)
|
||||
|
||||
return output.view(B, S_q, H, D_v)
|
||||
|
||||
|
||||
class AITerImpl(AttentionImpl):
|
||||
"""
|
||||
Implementation of attention using AITemplate.
|
||||
@@ -335,7 +142,7 @@ class AITerImpl(AttentionImpl):
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Performs attention using one of:
|
||||
- _mla_prefill_ps_attention (FP8, SGLANG_DIFFUSION_AITER_FP8_ATTN=1)
|
||||
- _fmha_fp8_prefill_attention (FP8, SGLANG_DIFFUSION_AITER_FP8_ATTN=1 when eligible)
|
||||
- flash_attn_func (BF16, default or FP8 fallback for unsupported shapes)
|
||||
|
||||
Args:
|
||||
@@ -357,8 +164,14 @@ class AITerImpl(AttentionImpl):
|
||||
one = torch.tensor(1.0, dtype=torch.float32, device=query.device)
|
||||
q_scale = k_scale = v_scale = one
|
||||
|
||||
if _can_use_mla_prefill(v_fp8.shape[-1], q_fp8.shape[2]):
|
||||
return _mla_prefill_ps_attention(
|
||||
d_q = q_fp8.shape[-1]
|
||||
d_k = k_fp8.shape[-1]
|
||||
d_v = v_fp8.shape[-1]
|
||||
h_q = q_fp8.shape[2]
|
||||
h_kv = k_fp8.shape[2]
|
||||
|
||||
if _can_use_fmha_fp8_prefill(d_q, d_k, d_v, h_q, h_kv):
|
||||
return _fmha_fp8_prefill_attention(
|
||||
q_fp8,
|
||||
k_fp8,
|
||||
v_fp8,
|
||||
@@ -370,13 +183,15 @@ class AITerImpl(AttentionImpl):
|
||||
)
|
||||
|
||||
logger.warning_once(
|
||||
"FP8 MLA prefill kernel unsupported "
|
||||
"(need gfx950, v_head_dim=%d, num_heads divisible by %d; "
|
||||
"got v_head_dim=%d, num_heads=%d). Falling back to BF16.",
|
||||
_MLA_PREFILL_V_HEAD_DIM,
|
||||
_MLA_PREFILL_HEAD_TILE,
|
||||
v_fp8.shape[-1],
|
||||
q_fp8.shape[2],
|
||||
"FP8 FMHA prefill unsupported for this shape (need gfx950-class AITER, "
|
||||
"full MHA, q/k/v head_dim=%d; got q=%d, k=%d, v=%d, num_heads=%d, "
|
||||
"num_kv_heads=%d). Falling back to BF16.",
|
||||
_FMHA_FP8_HEAD_DIM,
|
||||
d_q,
|
||||
d_k,
|
||||
d_v,
|
||||
h_q,
|
||||
h_kv,
|
||||
)
|
||||
|
||||
# BF16 path
|
||||
|
||||
Reference in New Issue
Block a user