diff --git a/python/sglang/srt/arg_groups/model_overrides/kimi_k3.py b/python/sglang/srt/arg_groups/model_overrides/kimi_k3.py index e8e3b58ad..fe2b361d6 100644 --- a/python/sglang/srt/arg_groups/model_overrides/kimi_k3.py +++ b/python/sglang/srt/arg_groups/model_overrides/kimi_k3.py @@ -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: diff --git a/python/sglang/srt/layers/attention/aiter_backend.py b/python/sglang/srt/layers/attention/aiter_backend.py index 7abaa12ed..7de8c09fc 100755 --- a/python/sglang/srt/layers/attention/aiter_backend.py +++ b/python/sglang/srt/layers/attention/aiter_backend.py @@ -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: diff --git a/python/sglang/srt/layers/attention/aiter_mla_gluon.py b/python/sglang/srt/layers/attention/aiter_mla_gluon.py index c2e9c524d..ee2f1aad3 100644 --- a/python/sglang/srt/layers/attention/aiter_mla_gluon.py +++ b/python/sglang/srt/layers/attention/aiter_mla_gluon.py @@ -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 diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index 8046af61b..33f76e2a6 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -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 diff --git a/python/sglang/srt/layers/dcp/comm.py b/python/sglang/srt/layers/dcp/comm.py index 8f2e1ef20..41e5ce284 100644 --- a/python/sglang/srt/layers/dcp/comm.py +++ b/python/sglang/srt/layers/dcp/comm.py @@ -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[ diff --git a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py index a730a8013..3ceab0d87 100644 --- a/python/sglang/srt/models/deepseek_common/attention_backend_handler.py +++ b/python/sglang/srt/models/deepseek_common/attention_backend_handler.py @@ -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 diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha_rocm.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha_rocm.py index 6b455d1d6..803df1dee 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha_rocm.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mha_rocm.py @@ -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 ) diff --git a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py index 1525f6c80..ca39935ab 100644 --- a/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py +++ b/python/sglang/srt/models/deepseek_common/attention_forward_methods/forward_mla.py @@ -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: 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 b62010ad2..91b3bc929 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 @@ -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(