[AMD] support qlen>1 for aiter gluon path for Kimi K3 (#37601)

This commit is contained in:
billishyahao
2026-09-06 19:25:58 -07:00
committed by GitHub
parent b9b39f222d
commit 1d5d85260c
4 changed files with 430 additions and 218 deletions
@@ -1467,10 +1467,16 @@ class GroupCoordinator:
# Bypass the function if we are using only 1 GPU.
if self.world_size == 1:
return input_
# Broadcast.
torch.distributed.broadcast(
input_, src=self.ranks[src], group=self.device_group
)
# Always use pynccl to avoid capturing hip graph failure on torch
# version smaller than or equal to 2.11
if is_hip() and self.pynccl_comm is not None and not self.pynccl_comm.disabled:
self.pynccl_comm.broadcast(input_, src=src)
else:
# Broadcast.
torch.distributed.broadcast(
input_, src=self.ranks[src], group=self.device_group
)
return input_
def broadcast_object(self, obj: Optional[Any] = None, src: int = 0):
@@ -1018,32 +1018,25 @@ class AiterAttnBackend(AttentionBackend):
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)
q = 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
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,
):
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,
return mla_gluon_decode(
q=q,
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,
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),
qlen=max_q_len,
)
if gluon_out is not None:
return gluon_out
work_metadata = self.forward_metadata.work_metadata
work_indptr = self.forward_metadata.work_indptr
@@ -1054,7 +1047,7 @@ class AiterAttnBackend(AttentionBackend):
num_kv_splits = self.forward_metadata.num_kv_splits
return self._mla_decode_fwd_with_head_pad(
q_mla,
q,
k_buffer.view(-1, 1, 1, layer.qk_head_dim),
layer,
qo_indptr=self.forward_metadata.qo_indptr,
@@ -2471,6 +2464,25 @@ class AiterAttnBackend(AttentionBackend):
)
return o
elif forward_batch.forward_mode.is_target_verify():
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,
):
return mla_gluon_decode(
q=q.view(-1, layer.tp_q_head_num, layer.qk_head_dim),
k_buffer=K_Buffer,
layer=layer,
kv_indices=self.forward_metadata.kv_indices,
kv_indptr=self.forward_metadata.kv_indptr,
sm_scale=layer.scaling,
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
),
qlen=self.forward_metadata.max_q_len or 1,
)
work_metadata = self.forward_metadata.work_metadata
work_indptr = self.forward_metadata.work_indptr
work_info_set = self.forward_metadata.work_info_set
@@ -2,7 +2,7 @@
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.
Gluon is unavailable.
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
@@ -11,9 +11,9 @@ is not pinned at build time.
from __future__ import annotations
import functools
import inspect
import logging
from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional
import torch
@@ -26,133 +26,49 @@ if TYPE_CHECKING:
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:
@functools.lru_cache(maxsize=1)
def _gluon_fn():
"""aiter mla gluon entry point return None if disabled"""
if not envs.SGLANG_AITER_MLA_GLUON.get():
logger.info("aiter mla gluon is disabled manually.")
return None
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)})"
from aiter.ops.triton.gluon.mla_gluon import mla_gluon
except ImportError as exc:
logger.info("aiter mla gluon import error message: %s", exc)
return None
# mla_gluon builds its shared layouts with cga_layout, added in Triton 3.7;
# older Triton fails at compile time with an opaque error instead.
if "cga_layout" not in inspect.signature(gl.PaddedSharedLayout).parameters:
logger.info(
"aiter mla gluon is disabled due to triton %s has no Gluon cga_layout (need >= 3.7)",
getattr(triton, "__version__", "unknown"),
)
return None
return mla_gluon
_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,
def log_mla_gluon_capability(log: logging.Logger | None = None) -> None:
"""Report whether Gluon MLA decode is valid; the reason is logged by _gluon_fn."""
ready = _gluon_fn() is not None
(log or logger).info(
"aiter mla gluon is %s",
"enabled" if ready else "disabled; falling back to zero-pad mla_decode_fwd",
)
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 prefer_mla_gluon_decode(
*, head_pad_mode: str, num_head: int, kv_cache_dtype: torch.dtype
) -> bool:
return (
head_pad_mode == "zero"
and num_head == 12
and kv_cache_dtype == fp8_dtype
and _gluon_fn() is not None
)
def mla_gluon_decode(
@@ -162,89 +78,46 @@ def mla_gluon_decode(
layer: RadixAttention,
kv_indices: torch.Tensor,
kv_indptr: torch.Tensor,
seq_lens: torch.Tensor,
sm_scale: float,
min_kv_seq_len: int,
kv_scale: float = 1.0,
min_kv_seq_len: Optional[int] = None,
qlen: int = 1,
) -> 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.
"""Run Gluon MLA decode for fused Q [num_tokens, H, 576] and MLA KV pool.
Returns [num_tokens, H, v_head_dim], or None when Gluon is unavailable.
"""
if not mla_gluon_available():
mla_gluon = _gluon_fn()
if mla_gluon is None:
return None
batch_size = q.shape[0]
num_head = layer.tp_q_head_num
kv_lora_rank = layer.v_head_dim
qk_rope_head_dim = layer.qk_head_dim - kv_lora_rank
batch_size = q.shape[0] // qlen
q_nope, q_pe = torch.split(q, [kv_lora_rank, qk_rope_head_dim], dim=-1)
if qlen > 1:
# Splitting the leading dim is a stride change, so these stay views of
# the non-contiguous torch.split outputs.
q_nope = q_nope.view(batch_size, qlen, num_head, kv_lora_rank)
q_pe = q_pe.view(batch_size, qlen, num_head, qk_rope_head_dim)
o = q.new_empty((batch_size, qlen, num_head, kv_lora_rank))
else:
o = q.new_empty((batch_size, num_head, kv_lora_rank))
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
mla_gluon(
q_nope,
q_pe,
k_buffer.view(-1, layer.qk_head_dim),
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,
)
# Hand back the caller's flat [num_tokens, H, v] layout either way.
return o.flatten(0, 1) if qlen > 1 else o