diff --git a/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py b/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py index 7f69be843..024c4d63b 100755 --- a/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py +++ b/python/sglang/multimodal_gen/runtime/layers/attention/backends/aiter.py @@ -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