[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") @_register_for("KimiK3ForConditionalGeneration")
def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict: def _kimi_k3_overrides(server_args: Any, hf_config: Any) -> dict:
cfg = resolving_view(server_args) 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", decode_attention_backend="tokenspeed_mla",
kv_cache_dtype="fp8_e4m3", 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: else:
raise AssertionError( 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: 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, scatter_req_to_token_to_page_table_kernel,
) )
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend 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.layers.dp_attention import is_dp_attention_enabled
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode from sglang.srt.model_executor.forward_batch_info import ForwardBatch, ForwardMode
from sglang.srt.speculative.spec_utils import ( 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." "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 ( from sglang.kernels.ops.attention.utils import (
launch_reshape_and_cache_flash, launch_reshape_and_cache_flash,
pad_sequence_with_mask, pad_sequence_with_mask,
@@ -151,11 +155,16 @@ class ForwardMetadata:
swa_page_table: Optional[torch.Tensor] = None swa_page_table: Optional[torch.Tensor] = None
# full->SWA translated out_cache_loc (SWA KV-store write target) # full->SWA translated out_cache_loc (SWA KV-store write target)
swa_out_cache_loc: Optional[torch.Tensor] = None 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 _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 # 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. # 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}) _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.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 # Get v_head_dim based on model type
if self.use_mla: if self.use_mla:
# For MLA models, get v_head_dim from model config # For MLA models, get v_head_dim from model config
@@ -348,6 +359,7 @@ class AiterAttnBackend(AttentionBackend):
# corresponding ForwardBatch fields. # corresponding ForwardBatch fields.
self.req_to_token_pool = model_runner.req_to_token_pool self.req_to_token_pool = model_runner.req_to_token_pool
self.token_to_kv_pool = model_runner.token_to_kv_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 # sliding window attention. Resolve the SWA pool rather than reading it
# straight off the active pool: a frozen-KV MTP draft worker's active # 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, # _mla_decode_fwd_with_head_pad brings any count below 16 up to it,
# by repetition when it divides 16 and by tiling otherwise. # by repetition when it divides 16 and by tiling otherwise.
_pad_heads_to_16 = self.num_head < 16 _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"Aiter MLA supports num_head of 4, 8, 12, or multiples of 16 "
f"in [16, 128].\n" f"in [16, 128].\n"
f"Provided {self.num_head} number of heads.\n" f"Provided {self.num_head} number of heads.\n"
"Try adjusting tensor_parallel_size value, or run decode on " "Try adjusting tensor_parallel_size value, or run decode on "
"another backend (--decode-attention-backend)." "another backend (--decode-attention-backend)."
) )
self.num_head_padded = 16 if self.num_head < 16 else self.num_head self.num_head_padded = 16 if self.num_head < 16 else self.num_head
if self.num_head in _mla_low_head_repeat: if self.num_head in _mla_low_head_repeat:
self.head_pad_mode = "repeat" self.head_pad_mode = "repeat"
@@ -451,17 +469,22 @@ class AiterAttnBackend(AttentionBackend):
self.head_pad_mode = "none" self.head_pad_mode = "none"
self.head_repeat_factor = 1 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.enable_dp_attention = is_dp_attention_enabled()
self.qo_indptr_ = torch.zeros( self.qo_indptr_ = torch.zeros(
(max_bs + 1,), dtype=torch.int32, device=model_runner.device (max_bs + 1,), dtype=torch.int32, device=model_runner.device
) )
global _use_mla_ps_kernel, fast_mode, intra_batch_mode 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 # 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 should not use fake-nps (fast_mode = False, intra_batch_mode = True)
# it will cause gpu-fault or accuracy issue # it will cause gpu-fault or accuracy issue.
if self.num_head in (32, 64, 128): if self.mla_kernel_num_head_padded in (32, 64, 128):
fast_mode = True fast_mode = True
intra_batch_mode = False intra_batch_mode = False
@@ -473,8 +496,8 @@ class AiterAttnBackend(AttentionBackend):
# Native 16-head persist is slow on TP8; keep disabled unless zero-pad # 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. # (e.g. Kimi K3 h12 -> qh16) where persist ASM is the fast path.
if ( if (
(self.num_head_padded == 16 and self.head_pad_mode != "zero") (self.mla_kernel_num_head_padded == 16 and self.head_pad_mode != "zero")
or self.num_head_padded == 128 or self.mla_kernel_num_head_padded == 128
) and self.kv_cache_dtype is not fp8_dtype: ) and self.kv_cache_dtype is not fp8_dtype:
_use_mla_ps_kernel = False _use_mla_ps_kernel = False
fast_mode = False fast_mode = False
@@ -551,7 +574,7 @@ class AiterAttnBackend(AttentionBackend):
return "fp8_e4m3" return "fp8_e4m3"
def make_mla_decode_meta_data_buffer(self, max_seqlen_qo, batch_size): 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 dtype = self.kv_cache_dtype
if self.enable_dp_attention: if self.enable_dp_attention:
@@ -637,7 +660,7 @@ class AiterAttnBackend(AttentionBackend):
qo_indptr, qo_indptr,
kv_indptr, kv_indptr,
kv_last_page_len, kv_last_page_len,
self.num_head_padded // nhead_kv, self.mla_kernel_num_head_padded // nhead_kv,
nhead_kv, nhead_kv,
False, False,
work_metadata, work_metadata,
@@ -1084,6 +1107,80 @@ class AiterAttnBackend(AttentionBackend):
num_kv_splits=num_kv_splits, 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( def mla_fp8_prefill_attn(
self, self,
q: torch.Tensor, q: torch.Tensor,
@@ -1243,6 +1340,9 @@ class AiterAttnBackend(AttentionBackend):
) )
max_kv_len = forward_batch.seq_lens_cpu.max().item() 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 forward_batch.forward_mode.is_decode_or_idle():
if spec_info is None or forward_batch.forward_mode.is_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) kv_indptr[1 : bs + 1] = torch.cumsum(forward_batch.seq_lens, dim=0)
@@ -1261,6 +1361,20 @@ class AiterAttnBackend(AttentionBackend):
kv_indices, kv_indices,
self.req_to_token.stride(0), 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: else:
max_q_len = 1 max_q_len = 1
page_size = self.page_size page_size = self.page_size
@@ -1317,7 +1431,9 @@ class AiterAttnBackend(AttentionBackend):
kv_last_page_len = self.kv_last_page_len[:bs] kv_last_page_len = self.kv_last_page_len[:bs]
max_q_len = 1 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_metadata,
work_indptr, work_indptr,
@@ -1456,9 +1572,13 @@ class AiterAttnBackend(AttentionBackend):
elif forward_batch.forward_mode.is_target_verify(): elif forward_batch.forward_mode.is_target_verify():
if self.use_mla: if self.use_mla:
draft_num = spec_info.draft_token_num 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 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 = self.qo_indptr[: bs + 1]
qo_indptr[: bs + 1] = torch.arange( qo_indptr[: bs + 1] = torch.arange(
@@ -1486,8 +1606,26 @@ class AiterAttnBackend(AttentionBackend):
TOKEN_BLOCK_PARALLEL=num_token_blocks > 1, TOKEN_BLOCK_PARALLEL=num_token_blocks > 1,
) )
# if self.kv_cache_dtype == fp8_dtype: if self.dcp_world_size > 1:
if _use_mla_ps_kernel: 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 max_seqlen_qo = draft_num
( (
work_metadata, work_metadata,
@@ -1532,6 +1670,8 @@ class AiterAttnBackend(AttentionBackend):
reduce_partial_map=reduce_partial_map, reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits, num_kv_splits=num_kv_splits,
run_graph=False, run_graph=False,
local_kv_lens=local_kv_lens,
verify_token_table=verify_token_table,
) )
else: else:
draft_num = forward_batch.input_ids.shape[0] // bs draft_num = forward_batch.input_ids.shape[0] // bs
@@ -1712,6 +1852,132 @@ class AiterAttnBackend(AttentionBackend):
swa_out_cache_loc=swa_out_cache_loc, 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( def init_cuda_graph_state(
self, self,
max_bs: int, max_bs: int,
@@ -1739,6 +2005,29 @@ class AiterAttnBackend(AttentionBackend):
self.cuda_graph_kv_last_page_len = torch.ones( self.cuda_graph_kv_last_page_len = torch.ones(
max_bs, dtype=torch.int32, device=self.device 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: if kv_indices_buf is None:
max_num_blocks_per_seq = ( max_num_blocks_per_seq = (
self.max_context_len + self.page_size - 1 self.max_context_len + self.page_size - 1
@@ -1854,6 +2143,10 @@ class AiterAttnBackend(AttentionBackend):
reduce_final_map = None reduce_final_map = None
reduce_partial_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 swa_page_table = None
max_kv_len = ( max_kv_len = (
seq_lens_cpu.max().item() seq_lens_cpu.max().item()
@@ -1887,6 +2180,20 @@ class AiterAttnBackend(AttentionBackend):
kv_indices, kv_indices,
self.req_to_token.stride(0), 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: else:
max_q_len = 1 max_q_len = 1
kv_indices = self.cuda_graph_page_table 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] kv_last_page_len = self.cuda_graph_kv_last_page_len[:bs]
max_q_len = 1 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 num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data( self.make_mla_meta_data(
@@ -2008,7 +2317,11 @@ class AiterAttnBackend(AttentionBackend):
device=self.device, device=self.device,
) )
if self.use_mla: 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: else:
kv_lens = seq_lens kv_lens = seq_lens
kv_indptr = self.kv_indptr[: bs + 1] 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] 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: if self.use_mla:
max_q_len = self.num_draft_tokens 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 num_kv_splits = self.max_split_per_batch
self.make_mla_meta_data( self.make_mla_meta_data(
@@ -2082,6 +2420,8 @@ class AiterAttnBackend(AttentionBackend):
reduce_final_map=reduce_final_map, reduce_final_map=reduce_final_map,
reduce_partial_map=reduce_partial_map, reduce_partial_map=reduce_partial_map,
num_kv_splits=num_kv_splits, num_kv_splits=num_kv_splits,
local_kv_lens=local_kv_lens,
verify_token_table=verify_token_table,
) )
else: else:
max_q_len = verify_tokens_per_req max_q_len = verify_tokens_per_req
@@ -2419,7 +2759,16 @@ class AiterAttnBackend(AttentionBackend):
v_scale=v_descale, v_scale=v_descale,
) )
elif self.use_mla: 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): elif self._use_fused_fp8_kv_write(layer):
# FP8: fuse bf16->fp8 cast + paged write in one kernel. # FP8: fuse bf16->fp8 cast + paged write in one kernel.
k_cache, v_cache = self.token_to_kv_pool.get_kv_buffer( 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) V_Buffer = self.token_to_kv_pool.get_value_buffer(layer.layer_id)
kv_lora_rank = V_Buffer.shape[-1] kv_lora_rank = V_Buffer.shape[-1]
qk_rope_head_dim = K_Buffer.shape[-1] - kv_lora_rank 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 qk_nope_head_dim = k.shape[-1] - qk_rope_head_dim
assert len(q.shape) == 3 assert len(q.shape) == 3
assert len(k.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) extend_no_prefix = not any(forward_batch.extend_prefix_lens_cpu)
if forward_batch.mha_return_lse: if forward_batch.mha_return_lse:
return self._forward_extend_skip_prefix(q, k, v, layer) 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 kv_indices.shape[0] == 0 or extend_no_prefix:
if self.use_fp8_prefill_attn and self.head_pad_mode != "zero": if self.use_fp8_prefill_attn and self.head_pad_mode != "zero":
output = self.mla_fp8_prefill_attn( output = self.mla_fp8_prefill_attn(
@@ -3253,6 +3625,10 @@ class AiterAttnBackend(AttentionBackend):
) )
if self.use_mla: 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) o = self._forward_mla_decode(q, layer, forward_batch, k_descale)
return o.reshape(-1, layer.tp_q_head_num * layer.v_head_dim) return o.reshape(-1, layer.tp_q_head_num * layer.v_head_dim)
else: else:
@@ -14,7 +14,7 @@ from __future__ import annotations
import functools import functools
import inspect import inspect
import logging import logging
from typing import TYPE_CHECKING, Optional from typing import TYPE_CHECKING
import torch import torch
@@ -82,9 +82,12 @@ def mla_gluon_decode(
min_kv_seq_len: int, min_kv_seq_len: int,
kv_scale: float = 1.0, kv_scale: float = 1.0,
qlen: int = 1, 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. """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() mla_gluon = _gluon_fn()
if mla_gluon is None: if mla_gluon is None:
@@ -105,7 +108,11 @@ def mla_gluon_decode(
else: else:
o = q.new_empty((batch_size, num_head, kv_lora_rank)) 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_nope,
q_pe, q_pe,
k_buffer.view(-1, layer.qk_head_dim), k_buffer.view(-1, layer.qk_head_dim),
@@ -115,9 +122,15 @@ def mla_gluon_decode(
sm_scale, sm_scale,
k_pe=None, k_pe=None,
kv_pe_offset=kv_lora_rank, kv_pe_offset=kv_lora_rank,
use_2d_view=False, use_2d_view=use_2d_view,
kv_scale=kv_scale, kv_scale=kv_scale,
min_kv_seq_len=min_kv_seq_len, min_kv_seq_len=min_kv_seq_len,
**extra_kwargs,
) )
# Hand back the caller's flat [num_tokens, H, v] layout either way. # 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_cuda,
is_gfx95_supported, is_gfx95_supported,
is_gfx942_supported, is_gfx942_supported,
is_hip,
is_xpu, is_xpu,
next_power_of_2, next_power_of_2,
) )
_is_cuda = is_cuda() _is_cuda = is_cuda()
_is_hip = is_hip()
_is_gfx942 = is_gfx942_supported() _is_gfx942 = is_gfx942_supported()
_is_xpu = is_xpu() _is_xpu = is_xpu()
@@ -229,8 +231,19 @@ class TritonAttnBackend(AttentionBackend):
self.use_mla, self.use_mla,
self.use_verify_splitkv, self.use_verify_splitkv,
) )
self.dcp_size = get_parallel().attn_dcp_size # TODO: this logic should be fixed in non-hip platform
self.dcp_rank = get_parallel().attn_dcp_rank 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 = ( self.num_head = (
model_runner.model_config.get_max_num_attention_heads() model_runner.model_config.get_max_num_attention_heads()
// get_parallel().attn_tp_size // 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.distributed.parallel_state import GroupCoordinator
from sglang.srt.runtime_context import get_parallel, get_platform 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 from sglang.srt.utils.common import is_mnnvl_fabric_device
_is_hip = is_hip()
def _warn_deprecated_dcp_accessor(name: str, replacement: str) -> None: def _warn_deprecated_dcp_accessor(name: str, replacement: str) -> None:
warnings.warn( warnings.warn(
@@ -276,19 +279,21 @@ def all_gather_kv_cache_for_mla_extend(
k_nope, k_nope,
k_pe, k_pe,
): ):
cache_k_nope, cache_k_rope = token_to_kv_pool.get_mla_kv_buffer( # On hip, skip the all-gather when there is no cached prefix to avoid crash
attn_mqa, if not _is_hip or dcp_extend_prefix_lens_sum > 0:
dcp_local_prefix_kv_indices, cache_k_nope, cache_k_rope = token_to_kv_pool.get_mla_kv_buffer(
) attn_mqa,
extend_prefix_lens_cpu = torch.tensor(extend_prefix_lens_cpu) dcp_local_prefix_kv_indices,
# all gather kv cache into forward_batch.attn_dcp_metadata.dcp_kv_buffer )
gathered_kv = all_gather_kv_cache_for_dcp( extend_prefix_lens_cpu = torch.tensor(extend_prefix_lens_cpu)
cache_k_nope, # all gather kv cache into forward_batch.attn_dcp_metadata.dcp_kv_buffer
cache_k_rope, gathered_kv = all_gather_kv_cache_for_dcp(
extend_prefix_lens_cpu, cache_k_nope,
prefix_starts_cpu=torch.zeros_like(extend_prefix_lens_cpu), cache_k_rope,
) extend_prefix_lens_cpu,
dcp_kv_buffer[:dcp_extend_prefix_lens_sum] = gathered_kv 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 # copy local kv cache into forward_batch.attn_dcp_metadata.dcp_kv_buffer
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 forward_batch.forward_mode.is_extend_without_speculative():
if not _support_mha_one_shot(attn, forward_batch, "aiter"): if not _support_mha_one_shot(attn, forward_batch, "aiter"):
return AttnForwardMethod.MHA_CHUNKED_KV return AttnForwardMethod.MHA_CHUNKED_KV
if get_parallel().dcp_enabled:
return AttnForwardMethod.MHA_ONE_SHOT
return AttnForwardMethod.MHA return AttnForwardMethod.MHA
else: else:
return AttnForwardMethod.MLA return AttnForwardMethod.MLA
@@ -14,16 +14,12 @@ import torch
from sglang.kernels.ops.attention.utils import concat_and_cast_mha_k_triton 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.communicator import get_attn_tp_context
from sglang.srt.layers.dcp import ( from sglang.srt.layers.dcp import all_gather_kv_cache_for_mha_extend
all_gather_kv_cache_for_mha_extend,
filter_dcp_local_kv_indices,
)
from sglang.srt.layers.quantization.fp8_utils import ( from sglang.srt.layers.quantization.fp8_utils import (
materialize_bpreshuffle_fp8_scale_tuple, materialize_bpreshuffle_fp8_scale_tuple,
) )
from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.model_executor.forward_context import ( from sglang.srt.model_executor.forward_context import (
get_attn_backend,
get_token_to_kv_pool, get_token_to_kv_pool,
) )
from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import ( from sglang.srt.models.deepseek_common.attention_forward_methods.forward_mha import (
@@ -319,11 +315,6 @@ class DeepseekMHARocmForwardMixin:
forward_batch: ForwardBatch, forward_batch: ForwardBatch,
): ):
if _use_aiter_gfx95: 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( kv_a, k_pe = get_token_to_kv_pool().get_mla_kv_buffer(
self.attn_mha, kv_indices, dst_dtype 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: 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: if _is_cuda:
@@ -591,23 +591,33 @@ class DeepseekMLARocmForwardMixin:
q_nope_out = apply_kv_b_lora_q_correction(self, q_nope, q_nope_out) 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) fuse_rope_for_trtllm_mla = self._fuse_rope_for_trtllm_mla(forward_batch)
if (
self.rotary_emb is not None force_rope_for_aiter_dcp_decode = (
and (not fuse_rope_for_trtllm_mla) get_parallel().dcp_enabled
and (not self._skip_rope_for_dsa_tilelang_fused())
and (not self._skip_rope_for_aiter_fused_mla())
and ( and (
not _use_aiter forward_batch.forward_mode.is_decode()
or not _is_gfx95_supported or forward_batch.forward_mode.is_target_verify()
or self.use_dsa or forward_batch.forward_mode.is_draft_extend_v2()
# Non-fused, non-specialized attention backends (e.g. Triton) run )
# the cat path in forward_absorb_core and need RoPE applied here; and _use_aiter_gfx95
# only the aiter fused MLA path and the specialized MLA backends and self.current_attention_backend
# defer RoPE to their own kernels. not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS
or ( )
self.current_attention_backend if self.rotary_emb is not None and (
not in FORWARD_ABSORB_CORE_ATTENTION_BACKENDS force_rope_for_aiter_dcp_decode
and self.current_attention_backend != "aiter" 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, q_pe=q_pe,
) )
elif forward_batch.forward_mode.is_extend(): elif forward_batch.forward_mode.is_extend():
# for extend, gather kv # Assemble the full sequence into dcp_kv_buffer, which the
all_gather_kv_cache_for_mla_extend( # backend attends over instead of the sharded local cache.
get_token_to_kv_pool(), if (
self.attn_mqa, forward_batch.attn_dcp_metadata is not None
forward_batch.extend_prefix_lens_cpu, and forward_batch.attn_dcp_metadata.dcp_kv_buffer is not None
forward_batch.attn_dcp_metadata.dcp_local_prefix_kv_indices, ):
forward_batch.attn_dcp_metadata.dcp_extend_prefix_lens_sum, all_gather_kv_cache_for_mla_extend(
forward_batch.attn_dcp_metadata.dcp_kv_buffer, get_token_to_kv_pool(),
self.kv_lora_rank, self.attn_mqa,
k_nope, forward_batch.extend_prefix_lens_cpu,
k_pe, 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: else:
logger.warning( logger.warning(
f"not supported forward_mode {forward_batch.forward_mode}" f"not supported forward_mode {forward_batch.forward_mode}"
@@ -763,6 +778,38 @@ class DeepseekMLARocmForwardMixin:
else {} 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: else:
if self._skip_rope_for_aiter_fused_mla(): if self._skip_rope_for_aiter_fused_mla():
q, _, _, k = _fused_rope_cat_and_cache( q, _, _, k = _fused_rope_cat_and_cache(