[AMD][DCP 1/N] add dcp support for aiter backend (#34432)

Co-authored-by: HAI <hixiao@gmail.com>
This commit is contained in:
billishyahao
2026-09-11 10:35:51 -07:00
committed by GitHub
co-authored by HAI
parent f69d6fc28a
commit 833bce9df5
9 changed files with 549 additions and 78 deletions
@@ -42,6 +42,16 @@ def _require_kimi_k3_cutedsl_dcp_support() -> None:
)
def _require_kimi_k3_aiter_gluon_dcp_support() -> None:
from sglang.srt.layers.attention.aiter_mla_gluon import _gluon_fn
if _gluon_fn() is None:
raise RuntimeError(
"Kimi-K3 DCP with decode_attention_backend='aiter' requires the aiter "
"gluon mla kernel, which is unavailable. See above aborting reasons."
)
@_register_for("KimiK3ForConditionalGeneration")
def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args)
@@ -101,9 +111,23 @@ def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
decode_attention_backend="tokenspeed_mla",
kv_cache_dtype="fp8_e4m3",
)
elif decode_backend == "aiter":
_require_kimi_k3_aiter_gluon_dcp_support()
# Override prefill backend to aiter by default
# if users don't explicitly specify triton
prefill_ab = "triton" if prefill_backend == "triton" else "aiter"
logger.info(
"Kimi-K3 DCP uses aiter MLA decode: "
f"prefill={prefill_backend!r} -> {prefill_ab!r}, "
f"decode={decode_backend!r} -> 'aiter'."
)
overrides.update(
prefill_attention_backend=prefill_ab,
decode_attention_backend="aiter",
)
else:
raise AssertionError(
f"Decode attention backend for Kimi-K3 DCP must be 'cutedsl_mla' or 'tokenspeed_mla', got {decode_backend!r}."
f"Decode attention backend for Kimi-K3 DCP must be 'cutedsl_mla', 'tokenspeed_mla' or 'aiter', got {decode_backend!r}."
)
if cfg.dcp_replicate_q_proj is None:
@@ -26,6 +26,8 @@ from sglang.kernels.ops.kvcache.aiter_unified_attention import (
scatter_req_to_token_to_page_table_kernel,
)
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
from sglang.srt.layers.dcp import update_local_kv_lens_for_dcp
from sglang.srt.layers.dcp.planner import plan_dcp_decode_metadata
from sglang.srt.layers.dp_attention import is_dp_attention_enabled
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.spec_utils import (
@@ -67,6 +69,8 @@ except ImportError:
"aiter is AMD specific kernel library. Please make sure aiter is installed on your AMD device."
)
from sglang.kernels.ops.attention.dcp_kernels import create_mla_kv_page_table_for_dcp
from sglang.kernels.ops.attention.merge_state import merge_state_triton
from sglang.kernels.ops.attention.utils import (
launch_reshape_and_cache_flash,
pad_sequence_with_mask,
@@ -151,11 +155,16 @@ class ForwardMetadata:
swa_page_table: Optional[torch.Tensor] = None
# full->SWA translated out_cache_loc (SWA KV-store write target)
swa_out_cache_loc: Optional[torch.Tensor] = None
local_kv_lens: Optional[torch.Tensor] = None
verify_token_table: Optional[torch.Tensor] = None
_AITER_PARTITION_SIZE_ROCM = 256
_DCP_VERIFY_TABLE_COLS_PER_BLOCK = 128
# AITER's gfx950 FP8 FMHA ASM kernels only cover these GQA ratios. Other
# ratios (e.g. Qwen3.8-27B 24Q/4KV = 6) must not take the pertensor shortcut.
_AITER_FP8_ASM_GQA_RATIOS = frozenset({1, 2, 4, 8, 16})
@@ -264,6 +273,8 @@ class AiterAttnBackend(AttentionBackend):
self.use_mla = model_runner.model_config.attention_arch == AttentionArch.MLA
self.dcp_world_size = get_parallel().attn_dcp_size
# Get v_head_dim based on model type
if self.use_mla:
# For MLA models, get v_head_dim from model config
@@ -348,6 +359,7 @@ class AiterAttnBackend(AttentionBackend):
# corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_pool
self.kv_index_translator = model_runner.kv_index_translator
# sliding window attention. Resolve the SWA pool rather than reading it
# straight off the active pool: a frozen-KV MTP draft worker's active
@@ -433,13 +445,19 @@ class AiterAttnBackend(AttentionBackend):
# _mla_decode_fwd_with_head_pad brings any count below 16 up to it,
# 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, (
assert (
self.dcp_world_size > 1
or _valid_heads
or _pad_heads_to_16
or not may_run_mla_decode
), (
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"
@@ -451,17 +469,22 @@ class AiterAttnBackend(AttentionBackend):
self.head_pad_mode = "none"
self.head_repeat_factor = 1
_gathered_num_head = self.num_head * self.dcp_world_size
self.mla_kernel_num_head_padded = (
16 if _gathered_num_head < 16 else _gathered_num_head
)
self.enable_dp_attention = is_dp_attention_enabled()
self.qo_indptr_ = torch.zeros(
(max_bs + 1,), dtype=torch.int32, device=model_runner.device
)
global _use_mla_ps_kernel, fast_mode, intra_batch_mode
# current mla_decode_fwd only support fake-nps in self.num_head == 16
# current mla_decode_fwd only support fake-nps in num_head == 16
# so all num_head size does not use qh16 kernel to simulate
# it should not use fake-nps (fast_mode = False, intra_batch_mode = True)
# it will cause gpu-fault or accuracy issue
if self.num_head in (32, 64, 128):
# it will cause gpu-fault or accuracy issue.
if self.mla_kernel_num_head_padded in (32, 64, 128):
fast_mode = True
intra_batch_mode = False
@@ -473,8 +496,8 @@ class AiterAttnBackend(AttentionBackend):
# 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 and self.head_pad_mode != "zero")
or self.num_head_padded == 128
(self.mla_kernel_num_head_padded == 16 and self.head_pad_mode != "zero")
or self.mla_kernel_num_head_padded == 128
) and self.kv_cache_dtype is not fp8_dtype:
_use_mla_ps_kernel = False
fast_mode = False
@@ -551,7 +574,7 @@ class AiterAttnBackend(AttentionBackend):
return "fp8_e4m3"
def make_mla_decode_meta_data_buffer(self, max_seqlen_qo, batch_size):
nhead = self.num_head_padded
nhead = self.mla_kernel_num_head_padded
dtype = self.kv_cache_dtype
if self.enable_dp_attention:
@@ -637,7 +660,7 @@ class AiterAttnBackend(AttentionBackend):
qo_indptr,
kv_indptr,
kv_last_page_len,
self.num_head_padded // nhead_kv,
self.mla_kernel_num_head_padded // nhead_kv,
nhead_kv,
False,
work_metadata,
@@ -1084,6 +1107,80 @@ class AiterAttnBackend(AttentionBackend):
num_kv_splits=num_kv_splits,
)
def _get_dcp_graph_max_local_kv_len(self) -> int:
"""Static upper bound on this rank's shard, ceil(max_context_len / W)."""
w = max(self.dcp_world_size, 1)
return (self.max_context_len + w - 1) // w
def _forward_decode_dcp(self, q, k_buffer, layer, k_descale):
"""Attend this rank's KV shard for decode -> (out, natural-log lse)."""
fm = self.forward_metadata
bs = fm.kv_indptr.shape[0] - 1
num_heads = layer.tp_q_head_num # gathered heads = num_local_heads * dcp
out, lse = mla_gluon_decode(
q=q.view(bs, num_heads, layer.qk_head_dim),
k_buffer=k_buffer,
layer=layer,
kv_indices=fm.kv_indices,
kv_indptr=fm.kv_indptr[: bs + 1],
sm_scale=layer.scaling,
kv_scale=self._resolve_fp8_kv_scale_float(layer, k_descale),
min_kv_seq_len=1,
return_lse=True,
)
return out, lse.view(bs, num_heads)
def _forward_verify_dcp(self, q, k_window, layer, k_descale):
"""Attend the committed prefix and the verify window separately, then
merge -> (out, natural-log lse).
Splitting at the window boundary avoids the one thing the decode kernel
cannot do under DCP: mask on the GLOBAL position g(j) = j * W + r.
"""
fm = self.forward_metadata
q_len = fm.max_q_len
num_heads = layer.tp_q_head_num # gathered heads = num_local_heads * dcp
seqused_k = fm.local_kv_lens
n_rows = seqused_k.shape[0]
bs = n_rows // q_len
out_a, lse_a = mla_gluon_decode(
q=q.view(n_rows, num_heads, layer.qk_head_dim),
k_buffer=self.token_to_kv_pool.get_key_buffer(layer.layer_id),
layer=layer,
kv_indices=fm.verify_token_table,
kv_indptr=seqused_k,
sm_scale=layer.scaling,
kv_scale=self._resolve_fp8_kv_scale_float(layer, k_descale),
min_kv_seq_len=1,
return_lse=True,
use_2d_view=True,
)
lse_a = lse_a.view(n_rows, num_heads)
# The verify window arrives as `k_window`, computed this forward and
# identical on every rank, so only one rank attends it; the others
# return their prefix partial for the cross-rank merge.
if get_parallel().attn_dcp_rank != 0:
return out_a, lse_a
# The window latent is request-major and contiguous, so it IS the pool:
# row i of request b lives at b * q_len + i. mla_gluon's MTP mask at
# seq_len == qlen is exactly the dense causal window this needs.
out_b, lse_b = mla_gluon_decode(
q=q.view(n_rows, num_heads, layer.qk_head_dim),
k_buffer=k_window,
layer=layer,
kv_indices=torch.arange(n_rows, dtype=torch.int32, device=q.device),
kv_indptr=torch.arange(bs + 1, dtype=torch.int32, device=q.device) * q_len,
sm_scale=layer.scaling,
min_kv_seq_len=1,
qlen=q_len,
return_lse=True,
)
return merge_state_triton(out_a, lse_a, out_b, lse_b.view(n_rows, num_heads))
def mla_fp8_prefill_attn(
self,
q: torch.Tensor,
@@ -1243,6 +1340,9 @@ class AiterAttnBackend(AttentionBackend):
)
max_kv_len = forward_batch.seq_lens_cpu.max().item()
# dcp metadata
local_kv_lens = None
verify_token_table = None
if forward_batch.forward_mode.is_decode_or_idle():
if spec_info is None or forward_batch.forward_mode.is_idle():
kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
@@ -1261,6 +1361,20 @@ class AiterAttnBackend(AttentionBackend):
kv_indices,
self.req_to_token.stride(0),
)
if (
self.use_mla
and self.dcp_world_size > 1
and not forward_batch.forward_mode.is_idle()
):
kv_lens = forward_batch.seq_lens[:bs].to(torch.int32).clone()
self._plan_dcp_decode_metadata(
kv_indptr,
kv_indices,
kv_lens,
forward_batch.seq_lens_cpu,
bs,
)
else:
max_q_len = 1
page_size = self.page_size
@@ -1317,7 +1431,9 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len = self.kv_last_page_len[:bs]
max_q_len = 1
if _use_mla_ps_kernel:
# DCP decode runs the aiter MLA kernel (builds its own block-table
# metadata in forward_decode), so skip the persist metadata.
if _use_mla_ps_kernel and self.dcp_world_size <= 1:
(
work_metadata,
work_indptr,
@@ -1456,9 +1572,13 @@ class AiterAttnBackend(AttentionBackend):
elif forward_batch.forward_mode.is_target_verify():
if self.use_mla:
draft_num = spec_info.draft_token_num
kv_lens = forward_batch.seq_lens + draft_num
kv_lens_sum = forward_batch.seq_lens_sum + draft_num * bs
device = forward_batch.seq_lens.device
if self.dcp_world_size > 1:
kv_lens = forward_batch.seq_lens.to(torch.int32).clone()
kv_lens_sum = forward_batch.seq_lens_sum
else:
kv_lens = forward_batch.seq_lens + draft_num
kv_lens_sum = forward_batch.seq_lens_sum + draft_num * bs
qo_indptr = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange(
@@ -1486,8 +1606,26 @@ class AiterAttnBackend(AttentionBackend):
TOKEN_BLOCK_PARALLEL=num_token_blocks > 1,
)
# if self.kv_cache_dtype == fp8_dtype:
if _use_mla_ps_kernel:
if self.dcp_world_size > 1:
self._plan_dcp_decode_metadata(
kv_indptr,
kv_indices,
kv_lens,
None,
bs,
)
(
verify_token_table,
local_kv_lens,
) = self._build_dcp_verify_token_table(
kv_indptr,
forward_batch.req_pool_indices,
bs,
draft_num,
(max_kv_len + self.dcp_world_size - 1) // self.dcp_world_size,
)
if _use_mla_ps_kernel and self.dcp_world_size <= 1:
max_seqlen_qo = draft_num
(
work_metadata,
@@ -1532,6 +1670,8 @@ class AiterAttnBackend(AttentionBackend):
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
run_graph=False,
local_kv_lens=local_kv_lens,
verify_token_table=verify_token_table,
)
else:
draft_num = forward_batch.input_ids.shape[0] // bs
@@ -1712,6 +1852,132 @@ class AiterAttnBackend(AttentionBackend):
swa_out_cache_loc=swa_out_cache_loc,
)
def _plan_dcp_decode_metadata(
self,
kv_indptr: torch.Tensor,
kv_indices: torch.Tensor,
kv_lens_gpu: torch.Tensor,
seq_lens_cpu: Optional[torch.Tensor],
bs: int,
static_local_kv_lens_cpu: Optional[torch.Tensor] = None,
):
"""Localize kv_indptr / kv_indices to this rank's DCP shard, in place."""
if static_local_kv_lens_cpu is not None:
# The planner reads `kv_len_arr_cpu` only for (max, sum), never for
# the lengths it writes back, so a static upper bound yields the same
# metadata without the GPU->CPU sync.
total_local_len = plan_dcp_decode_metadata(
kv_lens_gpu,
kv_indptr,
kv_indices,
init_metadata_replay=True,
fast_decode_kwargs={"kv_len_arr_cpu": static_local_kv_lens_cpu},
bs=bs,
)
elif seq_lens_cpu is not None:
kv_len_arr_cpu = seq_lens_cpu[:bs].to(torch.int32).clone()
update_local_kv_lens_for_dcp(kv_len_arr_cpu)
total_local_len = plan_dcp_decode_metadata(
kv_lens_gpu,
kv_indptr,
kv_indices,
init_metadata_replay=True,
fast_decode_kwargs={"kv_len_arr_cpu": kv_len_arr_cpu},
bs=bs,
)
else:
total_local_len = plan_dcp_decode_metadata(
kv_lens_gpu,
kv_indptr,
kv_indices,
init_metadata_replay=False,
fast_decode_kwargs={},
bs=bs,
)
# The planner leaves the compacted ids WIDENED (see its docstring), and
# mla_gluon indexes the pool directly, so collapse them once per forward
# -- the same contract flashinfer_mla_backend.py follows.
translator = self.kv_index_translator
if total_local_len > 0 and translator.needs_read_translate:
valid = kv_indices[:total_local_len]
valid.copy_(translator.translate_dcp_read_ids(valid))
def _build_dcp_local_kv_lens(
self,
kv_indptr: torch.Tensor,
bs: int,
out_lens: Optional[torch.Tensor] = None,
):
"""This rank's shard length per request, in TOKENS (mla_gluon's
``cache_seqlens``). ``kv_indptr`` must already be localized.
"""
lens = (kv_indptr[1 : bs + 1] - kv_indptr[:bs]).to(torch.int32)
if out_lens is None:
return lens
out_lens.copy_(lens)
return out_lens
def _build_dcp_verify_token_table(
self,
kv_indptr: torch.Tensor,
req_pool_indices: torch.Tensor,
bs: int,
q_len: int,
max_local_kv_len: int,
out: Optional[torch.Tensor] = None,
out_lens: Optional[torch.Tensor] = None,
):
"""Token table + shard lengths for the prefix attention of DCP verify.
One row per query token, one column per TOKEN (mla_gluon fixes
PAGE_SIZE at 1). Rows of a request repeat that request's shard.
"""
local_kv_lens = self._build_dcp_local_kv_lens(kv_indptr, bs)
n_rows = bs * q_len
if out is None:
# The row stride below is a Triton constexpr, so quantize the eager
# width: every distinct value costs a JIT specialization.
out = local_kv_lens.new_empty(
(
n_rows,
triton.cdiv(max_local_kv_len, _DCP_VERIFY_TABLE_COLS_PER_BLOCK)
* _DCP_VERIFY_TABLE_COLS_PER_BLOCK,
)
)
num_cols = out.shape[1]
# Write each request's row 0 in place: the row stride handed to the
# kernel spans that request's whole block of q_len rows.
translator = self.kv_index_translator
v2p = translator.full_v2p_table
create_mla_kv_page_table_for_dcp[
(bs, triton.cdiv(num_cols, _DCP_VERIFY_TABLE_COLS_PER_BLOCK))
](
self.req_to_token,
req_pool_indices,
local_kv_lens,
out,
v2p,
self.req_to_token.stride(0),
q_len * num_cols,
translator.full_page_multiplier,
PHYSICAL_PAGE_SIZE=1,
DCP_SIZE=self.dcp_world_size,
DCP_RANK=get_parallel().attn_dcp_rank,
PAGES_PER_BLOCK=_DCP_VERIFY_TABLE_COLS_PER_BLOCK,
HAS_V2P=v2p is not None,
)
rows = out.view(bs, q_len, num_cols)
if q_len > 1:
# Source is row 0, destination rows 1.., so the copy never overlaps.
rows[:, 1:, :].copy_(rows[:, :1, :].expand(bs, q_len - 1, num_cols))
if out_lens is None:
out_lens = local_kv_lens.new_empty((n_rows,))
out_lens.view(bs, q_len).copy_(local_kv_lens.unsqueeze(1).expand(bs, q_len))
return out, out_lens
def init_cuda_graph_state(
self,
max_bs: int,
@@ -1739,6 +2005,29 @@ class AiterAttnBackend(AttentionBackend):
self.cuda_graph_kv_last_page_len = torch.ones(
max_bs, dtype=torch.int32, device=self.device
)
if self.use_mla and self.dcp_world_size > 1:
if self.num_draft_tokens:
# Target-verify flattens the window into single-token rows, so it
# needs max_bs * num_draft_tokens.
n_verify_rows = max_bs * self.num_draft_tokens
self.cuda_graph_verify_local_kv_lens = torch.zeros(
(n_verify_rows,), dtype=torch.int32, device=self.device
)
# Capture-stable token table, filled out-of-graph. One column per
# TOKEN, so the width is the worst-case shard ceil(ctx_len / W).
self.cuda_graph_verify_token_table = torch.zeros(
(n_verify_rows, self._get_dcp_graph_max_local_kv_len()),
dtype=torch.int32,
device=self.device,
)
# Static per-rank shard bound, one entry per request. Sizes the
# verify plan without a sync; see _plan_dcp_decode_metadata.
self.cuda_graph_dcp_static_local_kv_lens = torch.full(
(max_bs,),
self._get_dcp_graph_max_local_kv_len(),
dtype=torch.int32,
device="cpu",
)
if kv_indices_buf is None:
max_num_blocks_per_seq = (
self.max_context_len + self.page_size - 1
@@ -1854,6 +2143,10 @@ class AiterAttnBackend(AttentionBackend):
reduce_final_map = None
reduce_partial_map = None
# DCP metadata which will be populated for MLA decode when dcp enabled
local_kv_lens = None
verify_token_table = None
swa_page_table = None
max_kv_len = (
seq_lens_cpu.max().item()
@@ -1887,6 +2180,20 @@ class AiterAttnBackend(AttentionBackend):
kv_indices,
self.req_to_token.stride(0),
)
if (
self.use_mla
and self.dcp_world_size > 1
and not forward_mode.is_idle()
):
kv_lens = seq_lens[:bs].to(torch.int32).clone()
self._plan_dcp_decode_metadata(
kv_indptr,
kv_indices,
kv_lens,
seq_lens_cpu,
bs,
)
else:
max_q_len = 1
kv_indices = self.cuda_graph_page_table
@@ -1946,7 +2253,9 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = 1
if _use_mla_ps_kernel:
# DCP decode builds its own block-table metadata in
# forward_decode, so the persist metadata is unused here.
if _use_mla_ps_kernel and self.dcp_world_size <= 1:
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
@@ -2008,7 +2317,11 @@ class AiterAttnBackend(AttentionBackend):
device=self.device,
)
if self.use_mla:
kv_lens = seq_lens + self.num_draft_tokens
kv_lens = (
seq_lens
if self.dcp_world_size > 1
else seq_lens + self.num_draft_tokens
)
else:
kv_lens = seq_lens
kv_indptr = self.kv_indptr[: bs + 1]
@@ -2039,9 +2352,34 @@ class AiterAttnBackend(AttentionBackend):
)
kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
if self.use_mla and self.dcp_world_size > 1:
self._plan_dcp_decode_metadata(
kv_indptr,
kv_indices,
seq_lens[:bs].to(torch.int32).clone(),
None,
bs,
static_local_kv_lens_cpu=self.cuda_graph_dcp_static_local_kv_lens[
:bs
],
)
n_rows = bs * self.num_draft_tokens
(
verify_token_table,
local_kv_lens,
) = self._build_dcp_verify_token_table(
kv_indptr,
req_pool_indices,
bs,
self.num_draft_tokens,
self._get_dcp_graph_max_local_kv_len(),
out=self.cuda_graph_verify_token_table[:n_rows],
out_lens=self.cuda_graph_verify_local_kv_lens[:n_rows],
)
if self.use_mla:
max_q_len = self.num_draft_tokens
if _use_mla_ps_kernel:
if _use_mla_ps_kernel and self.dcp_world_size <= 1:
num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data(
@@ -2082,6 +2420,8 @@ class AiterAttnBackend(AttentionBackend):
reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits,
local_kv_lens=local_kv_lens,
verify_token_table=verify_token_table,
)
else:
max_q_len = verify_tokens_per_req
@@ -2419,7 +2759,16 @@ class AiterAttnBackend(AttentionBackend):
v_scale=v_descale,
)
elif self.use_mla:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
if self.dcp_world_size > 1:
kv_lora_rank = v.shape[-1]
self.token_to_kv_pool.set_mla_kv_buffer(
layer,
cache_loc,
k[..., :kv_lora_rank],
k[..., kv_lora_rank:],
)
else:
self.token_to_kv_pool.set_kv_buffer(layer, cache_loc, k, v)
elif self._use_fused_fp8_kv_write(layer):
# FP8: fuse bf16->fp8 cast + paged write in one kernel.
k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer(
@@ -2458,6 +2807,15 @@ class AiterAttnBackend(AttentionBackend):
V_Buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
kv_lora_rank = V_Buffer.shape[-1]
qk_rope_head_dim = K_Buffer.shape[-1] - kv_lora_rank
if (
forward_batch.forward_mode.is_target_verify()
and self.dcp_world_size > 1
):
# two-stage dcp verify, dispatched before the dims below: the
# model provides the per rank kvcache slices.
return self._forward_verify_dcp(q, k, layer, k_descale)
qk_nope_head_dim = k.shape[-1] - qk_rope_head_dim
assert len(q.shape) == 3
assert len(k.shape) == 3
@@ -2471,6 +2829,20 @@ class AiterAttnBackend(AttentionBackend):
extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
if forward_batch.mha_return_lse:
return self._forward_extend_skip_prefix(q, k, v, layer)
if self.dcp_world_size > 1:
if self.use_fp8_prefill_attn and self.head_pad_mode != "zero":
return self.mla_fp8_prefill_attn(q, k, v, layer)
return flash_attn_varlen_func(
q,
k,
v,
qo_indptr,
forward_batch.attn_dcp_metadata.dcp_kv_indptr,
max_q_len,
max_kv_len,
softmax_scale=layer.scaling,
causal=True,
)
if kv_indices.shape[0] == 0 or extend_no_prefix:
if self.use_fp8_prefill_attn and self.head_pad_mode != "zero":
output = self.mla_fp8_prefill_attn(
@@ -3253,6 +3625,10 @@ class AiterAttnBackend(AttentionBackend):
)
if self.use_mla:
if self.dcp_world_size > 1 and not forward_batch.forward_mode.is_idle():
k_buffer = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
return self._forward_decode_dcp(q, k_buffer, layer, k_descale)
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:
@@ -14,7 +14,7 @@ from __future__ import annotations
import functools
import inspect
import logging
from typing import TYPE_CHECKING, Optional
from typing import TYPE_CHECKING
import torch
@@ -82,9 +82,12 @@ def mla_gluon_decode(
min_kv_seq_len: int,
kv_scale: float = 1.0,
qlen: int = 1,
) -> Optional[torch.Tensor]:
use_2d_view: bool = False,
return_lse: bool = False,
):
"""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.
Returns [num_tokens, H, v_head_dim] (or ``(out, lse)`` when ``return_lse``),
or None when Gluon is unavailable.
"""
mla_gluon = _gluon_fn()
if mla_gluon is None:
@@ -105,7 +108,11 @@ def mla_gluon_decode(
else:
o = q.new_empty((batch_size, num_head, kv_lora_rank))
mla_gluon(
extra_kwargs = {}
if return_lse:
extra_kwargs["return_lse"] = True
result = mla_gluon(
q_nope,
q_pe,
k_buffer.view(-1, layer.qk_head_dim),
@@ -115,9 +122,15 @@ def mla_gluon_decode(
sm_scale,
k_pe=None,
kv_pe_offset=kv_lora_rank,
use_2d_view=False,
use_2d_view=use_2d_view,
kv_scale=kv_scale,
min_kv_seq_len=min_kv_seq_len,
**extra_kwargs,
)
# Hand back the caller's flat [num_tokens, H, v] layout either way.
return o.flatten(0, 1) if qlen > 1 else o
out = o.flatten(0, 1) if qlen > 1 else o
if not return_lse:
return out
# mla_gluon writes the output into `o` and returns it alongside the lse.
_, lse = result
return out, lse
@@ -52,11 +52,13 @@ from sglang.srt.utils import (
is_cuda,
is_gfx95_supported,
is_gfx942_supported,
is_hip,
is_xpu,
next_power_of_2,
)
_is_cuda = is_cuda()
_is_hip = is_hip()
_is_gfx942 = is_gfx942_supported()
_is_xpu = is_xpu()
@@ -229,8 +231,19 @@ class TritonAttnBackend(AttentionBackend):
self.use_mla,
self.use_verify_splitkv,
)
self.dcp_size = get_parallel().attn_dcp_size
self.dcp_rank = get_parallel().attn_dcp_rank
# TODO: this logic should be fixed in non-hip platform
self.is_hip_dspark_draft = (
_is_hip
and model_runner.is_draft_worker
and model_runner.spec_algorithm.is_dspark()
)
if self.is_hip_dspark_draft:
# Drafts never join the dcp group so we ignore it
self.dcp_size = 1
self.dcp_rank = 0
else:
self.dcp_size = get_parallel().attn_dcp_size
self.dcp_rank = get_parallel().attn_dcp_rank
self.num_head = (
model_runner.model_config.get_max_num_attention_heads()
// get_parallel().attn_tp_size
+18 -13
View File
@@ -37,8 +37,11 @@ from sglang.srt.distributed.device_communicators.pynccl_allocator import (
)
from sglang.srt.distributed.parallel_state import GroupCoordinator
from sglang.srt.runtime_context import get_parallel, get_platform
from sglang.srt.utils import is_hip
from sglang.srt.utils.common import is_mnnvl_fabric_device
_is_hip = is_hip()
def _warn_deprecated_dcp_accessor(name: str, replacement: str) -> None:
warnings.warn(
@@ -276,19 +279,21 @@ def all_gather_kv_cache_for_mla_extend(
k_nope,
k_pe,
):
cache_k_nope, cache_k_rope = token_to_kv_pool.get_mla_kv_buffer(
attn_mqa,
dcp_local_prefix_kv_indices,
)
extend_prefix_lens_cpu = torch.tensor(extend_prefix_lens_cpu)
# all gather kv cache into forward_batch.attn_dcp_metadata.dcp_kv_buffer
gathered_kv = all_gather_kv_cache_for_dcp(
cache_k_nope,
cache_k_rope,
extend_prefix_lens_cpu,
prefix_starts_cpu=torch.zeros_like(extend_prefix_lens_cpu),
)
dcp_kv_buffer[:dcp_extend_prefix_lens_sum] = gathered_kv
# On hip, skip the all-gather when there is no cached prefix to avoid crash
if not _is_hip or dcp_extend_prefix_lens_sum > 0:
cache_k_nope, cache_k_rope = token_to_kv_pool.get_mla_kv_buffer(
attn_mqa,
dcp_local_prefix_kv_indices,
)
extend_prefix_lens_cpu = torch.tensor(extend_prefix_lens_cpu)
# all gather kv cache into forward_batch.attn_dcp_metadata.dcp_kv_buffer
gathered_kv = all_gather_kv_cache_for_dcp(
cache_k_nope,
cache_k_rope,
extend_prefix_lens_cpu,
prefix_starts_cpu=torch.zeros_like(extend_prefix_lens_cpu),
)
dcp_kv_buffer[:dcp_extend_prefix_lens_sum] = gathered_kv
# copy local kv cache into forward_batch.attn_dcp_metadata.dcp_kv_buffer
dcp_kv_buffer[
@@ -196,6 +196,8 @@ def handle_attention_aiter(attn, forward_batch):
if forward_batch.forward_mode.is_extend_without_speculative():
if not _support_mha_one_shot(attn, forward_batch, "aiter"):
return AttnForwardMethod.MHA_CHUNKED_KV
if get_parallel().dcp_enabled:
return AttnForwardMethod.MHA_ONE_SHOT
return AttnForwardMethod.MHA
else:
return AttnForwardMethod.MLA
@@ -14,16 +14,12 @@ import torch
from sglang.kernels.ops.attention.utils import concat_and_cast_mha_k_triton
from sglang.srt.layers.communicator import get_attn_tp_context
from sglang.srt.layers.dcp import (
all_gather_kv_cache_for_mha_extend,
filter_dcp_local_kv_indices,
)
from sglang.srt.layers.dcp import all_gather_kv_cache_for_mha_extend
from sglang.srt.layers.quantization.fp8_utils import (
materialize_bpreshuffle_fp8_scale_tuple,
)
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import (
get_attn_backend,
get_token_to_kv_pool,
)
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import (
@@ -319,11 +315,6 @@ class DeepseekMHARocmForwardMixin:
forward_batch: ForwardBatch,
):
if _use_aiter_gfx95:
kv_indices = filter_dcp_local_kv_indices(kv_indices=kv_indices)
# Read door: the pool never translates, so the production site does.
kv_indices = get_attn_backend().kv_index_translator.translate_dcp_read_ids(
kv_indices
)
kv_a, k_pe = get_token_to_kv_pool().get_mla_kv_buffer(
self.attn_mha, kv_indices, dst_dtype
)
@@ -99,7 +99,7 @@ def is_dcp_mla_decode_phase(forward_batch: ForwardBatch) -> bool:
def is_mla_dcp_lse_base_on_e(attention_backend: Optional[str]) -> bool:
return attention_backend in {"flashmla", "cutedsl_mla"}
return attention_backend in {"flashmla", "cutedsl_mla", "aiter"}
if _is_cuda:
@@ -591,23 +591,33 @@ class DeepseekMLARocmForwardMixin:
q_nope_out = apply_kv_b_lora_q_correction(self, q_nope, q_nope_out)
fuse_rope_for_trtllm_mla = self._fuse_rope_for_trtllm_mla(forward_batch)
if (
self.rotary_emb is not None
and (not fuse_rope_for_trtllm_mla)
and (not self._skip_rope_for_dsa_tilelang_fused())
and (not self._skip_rope_for_aiter_fused_mla())
force_rope_for_aiter_dcp_decode = (
get_parallel().dcp_enabled
and (
not _use_aiter
or not _is_gfx95_supported
or self.use_dsa
# Non-fused, non-specialized attention backends (e.g. Triton) run
# the cat path in forward_absorb_core and need RoPE applied here;
# only the aiter fused MLA path and the specialized MLA backends
# defer RoPE to their own kernels.
or (
self.current_attention_backend
not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS
and self.current_attention_backend != "aiter"
forward_batch.forward_mode.is_decode()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend_v2()
)
and _use_aiter_gfx95
and self.current_attention_backend
not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS
)
if self.rotary_emb is not None and (
force_rope_for_aiter_dcp_decode
or (
(not fuse_rope_for_trtllm_mla)
and (not self._skip_rope_for_dsa_tilelang_fused())
and (not self._skip_rope_for_aiter_fused_mla())
and (
not _use_aiter
or not _is_gfx95_supported
or self.use_dsa
or (
self.current_attention_backend
not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS
and self.current_attention_backend != "aiter"
)
)
)
):
@@ -622,18 +632,23 @@ class DeepseekMLARocmForwardMixin:
q_pe=q_pe,
)
elif forward_batch.forward_mode.is_extend():
# for extend, gather kv
all_gather_kv_cache_for_mla_extend(
get_token_to_kv_pool(),
self.attn_mqa,
forward_batch.extend_prefix_lens_cpu,
forward_batch.attn_dcp_metadata.dcp_local_prefix_kv_indices,
forward_batch.attn_dcp_metadata.dcp_extend_prefix_lens_sum,
forward_batch.attn_dcp_metadata.dcp_kv_buffer,
self.kv_lora_rank,
k_nope,
k_pe,
)
# Assemble the full sequence into dcp_kv_buffer, which the
# backend attends over instead of the sharded local cache.
if (
forward_batch.attn_dcp_metadata is not None
and forward_batch.attn_dcp_metadata.dcp_kv_buffer is not None
):
all_gather_kv_cache_for_mla_extend(
get_token_to_kv_pool(),
self.attn_mqa,
forward_batch.extend_prefix_lens_cpu,
forward_batch.attn_dcp_metadata.dcp_local_prefix_kv_indices,
forward_batch.attn_dcp_metadata.dcp_extend_prefix_lens_sum,
forward_batch.attn_dcp_metadata.dcp_kv_buffer,
self.kv_lora_rank,
k_nope,
k_pe,
)
else:
logger.warning(
f"not supported forward_mode {forward_batch.forward_mode}"
@@ -763,6 +778,38 @@ class DeepseekMLARocmForwardMixin:
else {}
),
)
elif (
_use_aiter
and (
forward_batch.forward_mode.is_decode()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend_v2()
)
and get_parallel().dcp_enabled
):
q = torch.cat([q_nope_out, q_pe], dim=-1)
if llama_4_scaling is not None:
q[..., : self.kv_lora_rank] *= llama_4_scaling
get_token_to_kv_pool().set_mla_kv_buffer(
self.attn_mqa,
forward_batch.out_cache_loc,
k_nope,
k_pe,
)
if forward_batch.forward_mode.is_target_verify():
k_window = torch.cat([k_nope, k_pe], dim=-1)
v_window = k_nope
else:
k_window = None
v_window = None
attn_output, lse = self.attn_mqa_for_dcp_decode(
q,
k_window,
v_window,
forward_batch,
save_kv_cache=False,
**(dict(topk_indices=topk_indices) if topk_indices is not None else {}),
)
else:
if self._skip_rope_for_aiter_fused_mla():
q, _, _, k = _fused_rope_cat_and_cache(