diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index d42e6b60a..8da36a531 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -858,6 +858,10 @@ class Envs: SGLANG_ROCM_USE_MULTI_STREAM = EnvBool(False) SGLANG_HACK_FLASHMLA_BACKEND = EnvStr("tilelang") SGLANG_USE_AITER_FP8_PER_TOKEN = EnvBool(False) + # Route Kimi-K3-style h12 + fp8 MLA decode through aiter Triton Gluon when + # import and Triton cga_layout prerequisites hold. Set to 0 to force the + # zero-pad mla_decode_fwd fallback (benchmarking / emergency disable). + SGLANG_AITER_MLA_GLUON = EnvBool(True) # DSV4 Aiter flags SGLANG_OPT_USE_AITER_SILU_MUL = EnvBool(False) diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 42892e8cf..37672d57f 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -79,6 +79,11 @@ from sglang.kernels.ops.quantization.fp8_kernel import ( ) from sglang.srt.configs.model_config import AttentionArch from sglang.srt.environ import envs +from sglang.srt.layers.attention.aiter_mla_gluon import ( + log_mla_gluon_capability, + mla_gluon_decode, + prefer_mla_gluon_decode, +) from sglang.srt.layers.attention.aiter_utils import ( forward_decode_vectorized_5d, forward_extend_vectorized_5d, @@ -398,10 +403,13 @@ class AiterAttnBackend(AttentionBackend): self.forward_metadata: ForwardMetadata = None if self.use_mla: - _valid_heads = self.num_head in (4, 8) or ( - self.num_head % 16 == 0 and 16 <= self.num_head <= 128 + _mla_low_head_repeat = (4, 8) + _mla_low_head_zero_pad = (12,) + _valid_heads = ( + self.num_head in _mla_low_head_repeat + or self.num_head in _mla_low_head_zero_pad + or (self.num_head % 16 == 0 and 16 <= self.num_head <= 128) ) - may_run_mla_decode = self.may_run_mla_decode_kernel( decode_attention_backend=model_runner.decode_attention_backend_str, speculative_algorithm=get_spec().speculative_algorithm, @@ -411,13 +419,22 @@ class AiterAttnBackend(AttentionBackend): # by repetition when it divides 16 and by tiling otherwise. _pad_heads_to_16 = self.num_head < 16 assert _valid_heads or _pad_heads_to_16 or not may_run_mla_decode, ( - f"Aiter MLA supports num_head of 4, 8, or multiples of 16 " + f"Aiter MLA supports num_head of 4, 8, 12, or multiples of 16 " f"in [16, 128].\n" f"Provided {self.num_head} number of heads.\n" "Try adjusting tensor_parallel_size value, or run decode on " "another backend (--decode-attention-backend)." ) self.num_head_padded = 16 if self.num_head < 16 else self.num_head + if self.num_head in _mla_low_head_repeat: + self.head_pad_mode = "repeat" + self.head_repeat_factor = 16 // self.num_head + elif self.num_head in _mla_low_head_zero_pad: + self.head_pad_mode = "zero" + self.head_repeat_factor = 1 + else: + self.head_pad_mode = "none" + self.head_repeat_factor = 1 self.enable_dp_attention = is_dp_attention_enabled() self.qo_indptr_ = torch.zeros( @@ -438,12 +455,21 @@ class AiterAttnBackend(AttentionBackend): # only use mla_ps_kernel when fp8 kv_cache # for non-fp8 kv_cache on tp8, use non-persist kernel to avoid performance degradation # head_num=16 (tp8 perf issue), head_num=128 (unsupported, like tp1 or --enable-dp-attention with tp8-dp8) + # Native 16-head persist is slow on TP8; keep disabled unless zero-pad + # (e.g. Kimi K3 h12 -> qh16) where persist ASM is the fast path. if ( - self.num_head_padded == 16 or self.num_head_padded == 128 + (self.num_head_padded == 16 and self.head_pad_mode != "zero") + or self.num_head_padded == 128 ) and self.kv_cache_dtype is not fp8_dtype: _use_mla_ps_kernel = False fast_mode = False intra_batch_mode = False + # Zero-pad topology (h12->qh16): prefer Gluon decode over PS kernel. + if self.head_pad_mode == "zero" and self.kv_cache_dtype == fp8_dtype: + _use_mla_ps_kernel = False + fast_mode = False + intra_batch_mode = False + log_mla_gluon_capability(logger) self.max_split_per_batch = 32 if _use_mla_ps_kernel else None @@ -906,11 +932,15 @@ class AiterAttnBackend(AttentionBackend): **kwargs, ): """Wrap mla_decode_fwd with head-dimension padding for num_head < 16. - The kernel only accepts 4, 8, or a multiple of 16 heads; any other - count crashes it. So we repeat q's heads to reach 16, real heads first. + + repeat (4/8): tile q heads to 16, slice back to num_head. + zero (12): pad four zero-valued dummy heads to 16 (vLLM #50371 style). q / o must already be shaped (..., num_head, head_dim). """ - if self.num_head_padded != self.num_head: + num_head = layer.tp_q_head_num + if self.head_pad_mode == "repeat" or ( + self.head_pad_mode == "none" and self.num_head_padded != self.num_head + ): q_in = self.pad_heads(q, self.num_head_padded) o = q.new_empty( (q.shape[0], self.num_head_padded, layer.v_head_dim), @@ -918,13 +948,127 @@ class AiterAttnBackend(AttentionBackend): ) mla_decode_fwd(q_in, k_buffer_flat, o, **kwargs) return o[:, : self.num_head, :] - else: + if self.head_pad_mode == "zero": + q_in = q.new_zeros( + (q.shape[0], self.num_head_padded, q.shape[-1]), + dtype=q.dtype, + ) + q_in[:, :num_head, :] = q o = q.new_empty( - (q.shape[0], layer.tp_q_head_num, layer.v_head_dim), + (q.shape[0], self.num_head_padded, layer.v_head_dim), dtype=self.input_dtype, ) - mla_decode_fwd(q, k_buffer_flat, o, **kwargs) - return o + mla_decode_fwd(q_in, k_buffer_flat, o, **kwargs) + return o[:, :num_head, :] + o = q.new_empty( + (q.shape[0], num_head, layer.v_head_dim), + dtype=self.input_dtype, + ) + mla_decode_fwd(q, k_buffer_flat, o, **kwargs) + return o + + def _zero_pad_mla_q_heads( + self, q: torch.Tensor, layer: RadixAttention + ) -> torch.Tensor: + """Zero-pad q heads num_head -> num_head_padded (12 -> 16) for the + aiter MLA prefill kernels. Input/return are 3-D (T, H, D); the extra + heads compute garbage outputs that are sliced away after the kernel.""" + q3 = q.view(-1, layer.tp_q_head_num, layer.qk_head_dim) + q_pad = q3.new_zeros((q3.shape[0], self.num_head_padded, q3.shape[-1])) + q_pad[:, : layer.tp_q_head_num, :] = q3 + return q_pad + + def _resolve_fp8_kv_scale_float(self, layer: RadixAttention, k_descale) -> float: + cached = getattr(layer, "_aiter_kv_scale_float", None) + if cached is not None: + return cached + if k_descale is None: + val = 1.0 + elif isinstance(k_descale, torch.Tensor): + val = float(k_descale.item()) + else: + val = float(k_descale) + layer._aiter_kv_scale_float = val + return val + + def _resolve_mla_gluon_min_kv_seq_len(self, forward_batch: ForwardBatch) -> int: + try: + if torch.cuda.is_current_stream_capturing(): + return int(self.max_context_len) + except Exception: + pass + if forward_batch.seq_lens_cpu is not None: + return int(forward_batch.seq_lens_cpu.max()) + seq_lens = forward_batch.seq_lens + if seq_lens is None or seq_lens.numel() == 0: + return 1 + return int(seq_lens.max().item()) + + def _forward_mla_decode( + self, + q: torch.Tensor, + layer: RadixAttention, + forward_batch: ForwardBatch, + k_descale, + ): + k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id) + q_mla = q.view(-1, layer.tp_q_head_num, layer.qk_head_dim) + max_q_len = self.forward_metadata.max_q_len or 1 + + if ( + prefer_mla_gluon_decode( + head_pad_mode=getattr(self, "head_pad_mode", "none"), + num_head=getattr(self, "num_head", layer.tp_q_head_num), + kv_cache_dtype=self.kv_cache_dtype, + ) + and max_q_len == 1 + ): + kv_scale = self._resolve_fp8_kv_scale_float(layer, k_descale) + min_kv_seq_len = self._resolve_mla_gluon_min_kv_seq_len(forward_batch) + gluon_out = mla_gluon_decode( + q=q_mla, + k_buffer=k_buffer, + layer=layer, + kv_indices=self.forward_metadata.kv_indices, + kv_indptr=self.forward_metadata.kv_indptr, + seq_lens=forward_batch.seq_lens, + sm_scale=layer.scaling, + kv_scale=kv_scale, + min_kv_seq_len=min_kv_seq_len, + ) + if gluon_out is not None: + return gluon_out + + work_metadata = self.forward_metadata.work_metadata + work_indptr = self.forward_metadata.work_indptr + work_info_set = self.forward_metadata.work_info_set + reduce_indptr = self.forward_metadata.reduce_indptr + reduce_final_map = self.forward_metadata.reduce_final_map + reduce_partial_map = self.forward_metadata.reduce_partial_map + num_kv_splits = self.forward_metadata.num_kv_splits + + return self._mla_decode_fwd_with_head_pad( + q_mla, + k_buffer.view(-1, 1, 1, layer.qk_head_dim), + layer, + qo_indptr=self.forward_metadata.qo_indptr, + kv_indptr=self.forward_metadata.kv_indptr, + kv_indices=self.forward_metadata.kv_indices, + kv_last_page_lens=self.forward_metadata.kv_last_page_len, + max_seqlen_q=max_q_len, + sm_scale=layer.scaling, + logit_cap=layer.logit_cap, + work_meta_data=work_metadata, + work_indptr=work_indptr, + work_info_set=work_info_set, + reduce_indptr=reduce_indptr, + reduce_final_map=reduce_final_map, + reduce_partial_map=reduce_partial_map, + q_scale=k_descale, + kv_scale=k_descale, + intra_batch_mode=intra_batch_mode, + num_kv_splits=num_kv_splits, + ) def mla_fp8_prefill_attn( self, @@ -1364,8 +1508,8 @@ class AiterAttnBackend(AttentionBackend): run_graph=False, ) else: - draft_num = forward_batch.input_ids.shape[0] // bs bs = len(forward_batch.req_pool_indices) + draft_num = spec_info.draft_token_num if self._use_unified_verify: page_table, qo_indptr, max_q_len, swa_page_table = ( @@ -1460,7 +1604,9 @@ class AiterAttnBackend(AttentionBackend): reduce_partial_map = None fp8_prefill_kv_indices = None - if self.use_fp8_prefill_attn: + # fp8 PS-ASM prefill memory-faults on gfx950 for 12-head + # (zero-pad) models; keep it off and use flash-attn fallback. + if self.use_fp8_prefill_attn and self.head_pad_mode != "zero": tile_q = 256 qlen_granularity = tile_q // ( self.fp8_prefill_num_head // self.fp8_prefill_num_kv_head @@ -2185,7 +2331,7 @@ class AiterAttnBackend(AttentionBackend): ): extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu) if kv_indices.shape[0] == 0 or extend_no_prefix: - if self.use_fp8_prefill_attn: + if self.use_fp8_prefill_attn and self.head_pad_mode != "zero": output = self.mla_fp8_prefill_attn( q, k, @@ -2219,6 +2365,7 @@ class AiterAttnBackend(AttentionBackend): if ( self.use_fp8_prefill_attn + and self.head_pad_mode != "zero" and layer.kv_b_proj.weight.dtype == torch.uint8 ): # MXFP4 weights + FP8 prefill: fuse GEMM, nope/v split, and k_pe cat @@ -2258,7 +2405,7 @@ class AiterAttnBackend(AttentionBackend): == forward_batch.extend_seq_lens.shape ) - if self.use_fp8_prefill_attn: + if self.use_fp8_prefill_attn and self.head_pad_mode != "zero": return self.mla_fp8_prefill_attn(q, k, v, layer) else: return flash_attn_varlen_func( @@ -2274,17 +2421,30 @@ class AiterAttnBackend(AttentionBackend): ) else: - if layer.qk_head_dim != layer.v_head_dim: + if self.head_pad_mode == "zero": + # 12 heads/rank (Kimi-K3 TP8): zero-pad q to 16 heads, + # run the aiter MLA prefill kernel, slice 12 back. + q_in = self._zero_pad_mla_q_heads(q, layer) o = q.new_empty( - (q.shape[0], layer.tp_q_head_num * layer.v_head_dim) + (q.shape[0], self.num_head_padded, layer.v_head_dim) ) else: - o = torch.empty_like(q) + q_in = q.view(-1, layer.tp_q_head_num, layer.qk_head_dim) + if layer.qk_head_dim != layer.v_head_dim: + o = q.new_empty( + (q.shape[0], layer.tp_q_head_num * layer.v_head_dim) + ) + else: + o = torch.empty_like(q) mla_prefill_fwd( - q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), + q_in, K_Buffer.view(-1, 1, 1, layer.qk_head_dim), - o.view(-1, layer.tp_q_head_num, layer.v_head_dim), + ( + o.view(-1, self.num_head_padded, layer.v_head_dim) + if self.head_pad_mode == "zero" + else o.view(-1, layer.tp_q_head_num, layer.v_head_dim) + ), qo_indptr, kv_indptr, kv_indices, @@ -2294,6 +2454,12 @@ class AiterAttnBackend(AttentionBackend): layer.logit_cap, ) K_Buffer = K_Buffer.view(-1, layer.tp_k_head_num, layer.qk_head_dim) + if self.head_pad_mode == "zero": + return ( + o[:, : layer.tp_q_head_num, :] + .contiguous() + .view(q.shape[0], layer.tp_q_head_num * layer.v_head_dim) + ) return o elif forward_batch.forward_mode.is_target_verify(): work_metadata = self.forward_metadata.work_metadata @@ -2503,9 +2669,7 @@ class AiterAttnBackend(AttentionBackend): v=v_unified, out=o.view(-1, layer.tp_q_head_num, layer.v_head_dim), cu_seqlens_q=self.forward_metadata.qo_indptr, - seqused_k=( - forward_batch.seq_lens + self.forward_metadata.max_q_len - ), + seqused_k=forward_batch.seq_lens + self.num_draft_tokens, max_seqlen_q=self.forward_metadata.max_q_len, max_seqlen_k=max_kv_len, softmax_scale=layer.scaling, @@ -2718,7 +2882,7 @@ class AiterAttnBackend(AttentionBackend): ) if o.dtype != self.input_dtype: o = o.to(self.input_dtype) - return o.view(-1, layer.tp_q_head_num * layer.head_dim) + return o.view(-1, layer.tp_q_head_num * layer.v_head_dim) if self.kv_cache_is_vectorized_5d: return forward_extend_vectorized_5d( @@ -2878,40 +3042,8 @@ class AiterAttnBackend(AttentionBackend): ) if self.use_mla: - k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id) - - work_metadata = self.forward_metadata.work_metadata - work_indptr = self.forward_metadata.work_indptr - work_info_set = self.forward_metadata.work_info_set - - reduce_indptr = self.forward_metadata.reduce_indptr - reduce_final_map = self.forward_metadata.reduce_final_map - reduce_partial_map = self.forward_metadata.reduce_partial_map - - num_kv_splits = self.forward_metadata.num_kv_splits - - o = self._mla_decode_fwd_with_head_pad( - q.view(-1, layer.tp_q_head_num, layer.qk_head_dim), - k_buffer.view(-1, 1, 1, layer.qk_head_dim), - layer, - qo_indptr=self.forward_metadata.qo_indptr, - kv_indptr=self.forward_metadata.kv_indptr, - kv_indices=self.forward_metadata.kv_indices, - kv_last_page_lens=self.forward_metadata.kv_last_page_len, - max_seqlen_q=self.forward_metadata.max_q_len, - sm_scale=layer.scaling, - logit_cap=layer.logit_cap, - work_meta_data=work_metadata, - work_indptr=work_indptr, - work_info_set=work_info_set, - reduce_indptr=reduce_indptr, - reduce_final_map=reduce_final_map, - reduce_partial_map=reduce_partial_map, - q_scale=k_descale, - kv_scale=k_descale, - intra_batch_mode=intra_batch_mode, - num_kv_splits=num_kv_splits, - ) + o = self._forward_mla_decode(q, layer, forward_batch, k_descale) + return o.reshape(-1, layer.tp_q_head_num * layer.v_head_dim) else: self.logits_soft_cap = layer.logit_cap diff --git a/python/sglang/srt/layers/attention/aiter_mla_gluon.py b/python/sglang/srt/layers/attention/aiter_mla_gluon.py new file mode 100644 index 000000000..50200c8f8 --- /dev/null +++ b/python/sglang/srt/layers/attention/aiter_mla_gluon.py @@ -0,0 +1,250 @@ +"""Gluon MLA decode wrapper for low head-count MLA (e.g. Kimi K3 TP8: 12 heads/GPU). + +Uses aiter ``mla_gluon`` when import succeeds and Triton Gluon exposes ``cga_layout`` +(needs Triton >= 3.7). Falls back to the caller (zero-pad + ``mla_decode_fwd``) when +Gluon is unavailable or fails at runtime. + +Requires aiter ``main`` with ROCm/aiter #4480 (batch>1 ``bh16bn128``) and #4555 +(decode CUDA graph KV splits). SGLang probes import + Triton API only; aiter version +is not pinned at build time. +""" + +from __future__ import annotations + +import inspect +import logging +from dataclasses import dataclass +from typing import TYPE_CHECKING, Optional + +import torch + +from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype +from sglang.srt.environ import envs + +if TYPE_CHECKING: + from sglang.srt.layers.radix_attention import RadixAttention + +logger = logging.getLogger(__name__) + +_mla_gluon_fn = None +_mla_gluon_import_failed = False +_capability_cache: Optional[MlaGluonCapability] = None + + +@dataclass(frozen=True) +class MlaGluonCapability: + """Runtime probe of aiter/Triton Gluon prerequisites for h12 + FP8 decode.""" + + enabled_by_env: bool + import_ok: bool + triton_version: str + triton_cga_layout_ok: bool + ready: bool + summary: str + + def missing_for_ready(self) -> list[str]: + missing = [] + if not self.enabled_by_env: + missing.append("SGLANG_AITER_MLA_GLUON=0") + if not self.import_ok: + missing.append("aiter.ops.triton.gluon.mla_gluon import") + if not self.triton_cga_layout_ok: + missing.append( + f"Triton Gluon cga_layout (have {self.triton_version or 'unknown'}, need >= 3.7)" + ) + return missing + + +def _triton_version() -> str: + try: + import triton + + return getattr(triton, "__version__", "unknown") + except Exception: + return "missing" + + +def _triton_cga_layout_ok() -> bool: + try: + import triton.experimental.gluon.language as gl + + return "cga_layout" in inspect.signature(gl.PaddedSharedLayout).parameters + except Exception: + return False + + +def _gluon_runtime_ok() -> bool: + return mla_gluon_available() and _triton_cga_layout_ok() + + +def _mla_gluon_enabled() -> bool: + return envs.SGLANG_AITER_MLA_GLUON.get() + + +def probe_mla_gluon_capability(*, force_refresh: bool = False) -> MlaGluonCapability: + global _capability_cache + if _capability_cache is not None and not force_refresh: + return _capability_cache + + enabled = _mla_gluon_enabled() + triton_ver = _triton_version() + import_ok = mla_gluon_available() if enabled else False + cga_ok = _triton_cga_layout_ok() + ready = enabled and import_ok and cga_ok + + if ready: + summary = f"Gluon MLA h12+fp8 ready (Triton={triton_ver})" + else: + cap = MlaGluonCapability( + enabled_by_env=enabled, + import_ok=import_ok, + triton_version=triton_ver, + triton_cga_layout_ok=cga_ok, + ready=False, + summary="", + ) + missing = cap.missing_for_ready() + summary = ( + "Gluon MLA h12+fp8 disabled; fallback to zero-pad mla_decode_fwd " + f"({'; '.join(missing)})" + ) + + _capability_cache = MlaGluonCapability( + enabled_by_env=enabled, + import_ok=import_ok, + triton_version=triton_ver, + triton_cga_layout_ok=cga_ok, + ready=ready, + summary=summary, + ) + return _capability_cache + + +def log_mla_gluon_capability(log: logging.Logger | None = None) -> MlaGluonCapability: + cap = probe_mla_gluon_capability() + (log or logger).info(cap.summary) + if not cap.ready: + for item in cap.missing_for_ready(): + (log or logger).info(" missing: %s", item) + return cap + + +def _in_cuda_graph_capture() -> bool: + try: + return bool(torch.cuda.is_current_stream_capturing()) + except Exception: + return False + + +def mla_gluon_available() -> bool: + if not _mla_gluon_enabled(): + return False + global _mla_gluon_fn, _mla_gluon_import_failed + if _mla_gluon_import_failed: + return False + if _mla_gluon_fn is not None: + return True + try: + from aiter.ops.triton.gluon.mla_gluon import mla_gluon as fn + + _mla_gluon_fn = fn + return True + except ImportError: + _mla_gluon_import_failed = True + logger.warning("mla_gluon import failed; Gluon MLA decode disabled.") + return False + + +def mla_gluon_decode( + *, + q: torch.Tensor, + k_buffer: torch.Tensor, + layer: RadixAttention, + kv_indices: torch.Tensor, + kv_indptr: torch.Tensor, + seq_lens: torch.Tensor, + sm_scale: float, + kv_scale: float = 1.0, + min_kv_seq_len: Optional[int] = None, +) -> Optional[torch.Tensor]: + """Run Gluon MLA decode for fused Q [B, H, 576] and MLA KV pool. + + Returns output [B, H, v_head_dim] on success, or None to fall back. + + ``min_kv_seq_len`` must be supplied by the caller during CUDA graph capture + (no GPU->CPU sync from ``seq_lens``). For eager decode, omit it to derive + from ``seq_lens`` when safe. + """ + if not mla_gluon_available(): + return None + + batch_size = q.shape[0] + + kv_lora_rank = layer.v_head_dim + qk_rope_head_dim = layer.qk_head_dim - kv_lora_rank + q_nope, q_pe = torch.split(q, [kv_lora_rank, qk_rope_head_dim], dim=-1) + + o = q.new_empty((batch_size, layer.tp_q_head_num, kv_lora_rank)) + + kv_c = k_buffer.view(-1, layer.qk_head_dim) + if min_kv_seq_len is None: + if _in_cuda_graph_capture(): + logger.warning( + "mla_gluon_decode: min_kv_seq_len missing during CUDA graph capture" + ) + min_kv_seq_len = 1 + elif seq_lens.numel(): + min_kv_seq_len = int(seq_lens.max().item()) + else: + min_kv_seq_len = 1 + + try: + _mla_gluon_fn( + q_nope, + q_pe, + kv_c, + o, + kv_indices, + kv_indptr, + sm_scale, + k_pe=None, + kv_pe_offset=kv_lora_rank, + use_2d_view=False, + kv_scale=kv_scale, + min_kv_seq_len=min_kv_seq_len, + ) + return o + except Exception as exc: + logger.warning( + "mla_gluon decode failed (num_head=%s, kv_dtype=%s, batch=%s): %s; " + "falling back to zero-pad mla_decode_fwd", + layer.tp_q_head_num, + k_buffer.dtype, + batch_size, + exc, + ) + return None + + +def prefer_mla_gluon_decode( + *, head_pad_mode: str, num_head: int, kv_cache_dtype: torch.dtype +) -> bool: + """Route Kimi-style h12 zero-pad MLA decode through Gluon when FP8 KV holds. + + ``head_pad_mode == "zero"`` selects the legacy ``mla_decode_fwd`` padding + topology (N heads padded to 16). Gluon is only validated for ``num_head == 12`` + today; other zero-pad head counts must stay on zero-pad + ``mla_decode_fwd``. + """ + if not _mla_gluon_enabled(): + return False + if head_pad_mode == "zero" and num_head == 12 and kv_cache_dtype == fp8_dtype: + return _gluon_runtime_ok() + return False + + +def reset_mla_gluon_state_for_test() -> None: + """Test helper: clear import/probe caches.""" + global _mla_gluon_fn, _mla_gluon_import_failed, _capability_cache + _mla_gluon_fn = None + _mla_gluon_import_failed = False + _capability_cache = None diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py index 6cbe7e765..e7fe7f7a5 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla_rocm.py @@ -317,6 +317,13 @@ def _fused_rope_cat_and_cache( kv_cache_dtype = ( fp8_dtype if attn.kv_cache_dtype == "fp8_e4m3" else q_nope_out.dtype ) + # Gluon MLA decode (bh16bn128) requires bf16 Q; vLLM #50563. + q_out_dtype = ( + q_nope_out.dtype + if attn.kv_cache_dtype == "fp8_e4m3" + and attn.current_attention_backend == "aiter" + else kv_cache_dtype + ) return fused_qk_rope_cat_and_cache_mla( q_nope_out, q_pe, @@ -329,7 +336,7 @@ def _fused_rope_cat_and_cache( attn.rotary_emb.sin_cache, attn.attn_mqa.k_scale, attn.rotary_emb.is_neox_style, - q_out_dtype=kv_cache_dtype, + q_out_dtype=q_out_dtype, ) @@ -906,6 +913,8 @@ class DeepseekMLARocmForwardMixin: rotary_emb.cos_cache, so skipping the standalone rope there ends in AttributeError on None. Kimi-K3 has such layers. """ + # NoPE models (rotary_emb=None, e.g. Kimi-K3) have no rope for the + # fused kernel to apply; keep both prepare and core on the plain path. return ( _use_aiter_gfx95 and self.current_attention_backend == "aiter" diff --git a/test/registered/attention/test_mla_gluon_h12_fp8.py b/test/registered/attention/test_mla_gluon_h12_fp8.py new file mode 100644 index 000000000..5bbeedecc --- /dev/null +++ b/test/registered/attention/test_mla_gluon_h12_fp8.py @@ -0,0 +1,205 @@ +"""Unit tests for h12 + FP8 Gluon routing and zero-pad fallback dispatch. + +CPU-only mocks — no aiter/Triton/GPU required. +""" + +import unittest +from unittest import mock + +import torch + +from sglang.srt.environ import envs +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase + +register_cpu_ci(est_time=15, suite="base-a-test-cpu") + + +class TestMlaGluonCapability(CustomTestCase): + def setUp(self): + from sglang.srt.layers.attention.aiter_mla_gluon import ( + reset_mla_gluon_state_for_test, + ) + + reset_mla_gluon_state_for_test() + + def tearDown(self): + from sglang.srt.layers.attention.aiter_mla_gluon import ( + reset_mla_gluon_state_for_test, + ) + + reset_mla_gluon_state_for_test() + + def test_env_disable_not_ready(self): + from sglang.srt.layers.attention import aiter_mla_gluon as mod + + with envs.SGLANG_AITER_MLA_GLUON.override(False): + mod.reset_mla_gluon_state_for_test() + cap = mod.probe_mla_gluon_capability(force_refresh=True) + self.assertFalse(cap.ready) + self.assertIn("SGLANG_AITER_MLA_GLUON=0", cap.missing_for_ready()) + + @mock.patch( + "sglang.srt.layers.attention.aiter_mla_gluon._triton_cga_layout_ok", + return_value=True, + ) + @mock.patch( + "sglang.srt.layers.attention.aiter_mla_gluon._triton_version", + return_value="3.7.0", + ) + def test_ready_when_import_and_cga_ok(self, _ver, _cga): + from sglang.srt.layers.attention import aiter_mla_gluon as mod + + fake_fn = mock.Mock() + with mock.patch.dict( + "sys.modules", + { + "aiter": mock.MagicMock(), + "aiter.ops": mock.MagicMock(), + "aiter.ops.triton": mock.MagicMock(), + "aiter.ops.triton.gluon": mock.MagicMock(), + "aiter.ops.triton.gluon.mla_gluon": mock.MagicMock(mla_gluon=fake_fn), + }, + ): + mod.reset_mla_gluon_state_for_test() + cap = mod.probe_mla_gluon_capability(force_refresh=True) + self.assertTrue(cap.ready) + self.assertIn("3.7.0", cap.summary) + + @mock.patch( + "sglang.srt.layers.attention.aiter_mla_gluon._triton_cga_layout_ok", + return_value=False, + ) + @mock.patch( + "sglang.srt.layers.attention.aiter_mla_gluon.mla_gluon_available", + return_value=True, + ) + def test_prefer_false_when_cga_missing(self, _avail, _cga): + from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype + from sglang.srt.layers.attention.aiter_mla_gluon import prefer_mla_gluon_decode + + self.assertFalse( + prefer_mla_gluon_decode( + head_pad_mode="zero", num_head=12, kv_cache_dtype=fp8_dtype + ) + ) + + @mock.patch( + "sglang.srt.layers.attention.aiter_mla_gluon._gluon_runtime_ok", + return_value=True, + ) + def test_prefer_false_when_zero_pad_but_not_h12(self, _ok): + from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype + from sglang.srt.layers.attention.aiter_mla_gluon import prefer_mla_gluon_decode + + self.assertFalse( + prefer_mla_gluon_decode( + head_pad_mode="zero", num_head=10, kv_cache_dtype=fp8_dtype + ) + ) + + +class TestMlaGluonDecodeFallback(CustomTestCase): + """Verify _forward_mla_decode uses zero-pad path when Gluon is off or fails.""" + + def _make_backend(self): + from sglang.kernels.ops.quantization.fp8_kernel import fp8_dtype + from sglang.srt.layers.attention.aiter_backend import AiterAttnBackend + + be = AiterAttnBackend.__new__(AiterAttnBackend) + be.num_head = 12 + be.kv_cache_dtype = fp8_dtype + be.head_pad_mode = "zero" + be.num_head_padded = 16 + be.forward_metadata = mock.Mock( + max_q_len=1, + kv_indices=torch.zeros(4, dtype=torch.int32), + kv_indptr=torch.tensor([0, 1, 2, 3, 4], dtype=torch.int32), + kv_last_page_len=torch.ones(4, dtype=torch.int32), + qo_indptr=torch.arange(5, dtype=torch.int32), + work_metadata=None, + work_indptr=None, + work_info_set=None, + reduce_indptr=None, + reduce_final_map=None, + reduce_partial_map=None, + num_kv_splits=None, + ) + be.token_to_kv_pool = mock.Mock( + get_key_buffer=lambda _lid: torch.zeros(8, 576, dtype=fp8_dtype) + ) + be._resolve_fp8_kv_scale_float = mock.Mock(return_value=1.0) + be._resolve_mla_gluon_min_kv_seq_len = mock.Mock(return_value=128) + be._mla_decode_fwd_with_head_pad = mock.Mock( + return_value=torch.zeros(4, 12, 512) + ) + return be + + def _make_layer(self): + layer = mock.Mock() + layer.tp_q_head_num = 12 + layer.qk_head_dim = 576 + layer.v_head_dim = 512 + layer.scaling = 0.125 + layer.logit_cap = 0.0 + layer.layer_id = 0 + return layer + + @mock.patch( + "sglang.srt.layers.attention.aiter_backend.prefer_mla_gluon_decode", + return_value=False, + ) + @mock.patch("sglang.srt.layers.attention.aiter_backend.mla_gluon_decode") + def test_skips_gluon_when_disabled(self, mock_gluon, _prefer): + be = self._make_backend() + layer = self._make_layer() + q = torch.zeros(4, 12, 576, dtype=torch.bfloat16) + fb = mock.Mock(seq_lens=torch.tensor([128, 128, 128, 128])) + + out = be._forward_mla_decode(q, layer, fb, k_descale=1.0) + + mock_gluon.assert_not_called() + be._mla_decode_fwd_with_head_pad.assert_called_once() + self.assertIs(out, be._mla_decode_fwd_with_head_pad.return_value) + + @mock.patch( + "sglang.srt.layers.attention.aiter_backend.prefer_mla_gluon_decode", + return_value=True, + ) + @mock.patch( + "sglang.srt.layers.attention.aiter_backend.mla_gluon_decode", + return_value=None, + ) + def test_falls_back_when_gluon_returns_none(self, mock_gluon, _prefer): + be = self._make_backend() + layer = self._make_layer() + q = torch.zeros(4, 12, 576, dtype=torch.bfloat16) + fb = mock.Mock(seq_lens=torch.tensor([128, 128, 128, 128])) + + be._forward_mla_decode(q, layer, fb, k_descale=1.0) + + mock_gluon.assert_called_once() + be._mla_decode_fwd_with_head_pad.assert_called_once() + + @mock.patch( + "sglang.srt.layers.attention.aiter_backend.prefer_mla_gluon_decode", + return_value=True, + ) + @mock.patch("sglang.srt.layers.attention.aiter_backend.mla_gluon_decode") + def test_uses_gluon_output_when_ok(self, mock_gluon, _prefer): + gluon_out = torch.ones(4, 12, 512) + mock_gluon.return_value = gluon_out + be = self._make_backend() + layer = self._make_layer() + q = torch.zeros(4, 12, 576, dtype=torch.bfloat16) + fb = mock.Mock(seq_lens=torch.tensor([128, 128, 128, 128])) + + out = be._forward_mla_decode(q, layer, fb, k_descale=1.0) + + mock_gluon.assert_called_once() + be._mla_decode_fwd_with_head_pad.assert_not_called() + self.assertIs(out, gluon_out) + + +if __name__ == "__main__": + unittest.main()