Support DCP for Kimi Linear model (#32612)

Co-authored-by: Julien Lin <jullin@nvidia.com>
Co-authored-by: kpham-sgl <khoa.pham@radixark.ai>
This commit is contained in:
Baizhou Zhang
2026-07-28 22:59:58 -07:00
committed by GitHub
co-authored by Julien Lin kpham-sgl
parent c4fc241fd3
commit ef6c07008b
17 changed files with 1331 additions and 86 deletions
@@ -76,6 +76,40 @@ def create_triton_kv_indices_for_dcp_triton(
# KV-index build (PR #14194, MLA): global prefix+extend layout for the
# all-gathered dcp_kv_buffer, plus the per-rank shard/compact kernel.
# ---------------------------------------------------------------------------
@triton.jit
def create_mla_kv_page_table_for_dcp(
req_to_token_ptr,
req_pool_indices_ptr,
local_seq_lens_ptr,
block_kv_indices_ptr,
req_to_token_stride: tl.constexpr,
block_table_stride: tl.constexpr,
PHYSICAL_PAGE_SIZE: tl.constexpr,
DCP_SIZE: tl.constexpr,
DCP_RANK: tl.constexpr,
PAGES_PER_BLOCK: tl.constexpr,
):
req = tl.program_id(0)
page_block = tl.program_id(1)
page_offsets = page_block * PAGES_PER_BLOCK + tl.arange(0, PAGES_PER_BLOCK)
local_len = tl.load(local_seq_lens_ptr + req)
local_pages = tl.cdiv(local_len, PHYSICAL_PAGE_SIZE)
mask = page_offsets < local_pages
global_positions = DCP_RANK + page_offsets * PHYSICAL_PAGE_SIZE * DCP_SIZE
req_pool_index = tl.load(req_pool_indices_ptr + req)
virtual_locs = tl.load(
req_to_token_ptr + req_pool_index * req_to_token_stride + global_positions,
mask=mask,
other=0,
)
physical_pages = virtual_locs // DCP_SIZE // PHYSICAL_PAGE_SIZE
tl.store(
block_kv_indices_ptr + req * block_table_stride + page_offsets,
physical_pages,
mask=mask,
)
@triton.jit
def create_dcp_kv_indices(
kv_indptr,
@@ -89,9 +89,9 @@ def create_tokenspeed_mla_backend(runner):
def create_cutedsl_mla_backend(runner):
if not runner.use_mla_backend:
raise ValueError("cutedsl_mla backend can only be used with MLA models.")
from sglang.srt.layers.attention.trtllm_mla_backend import TRTLLMMLABackend
from sglang.srt.layers.attention.cutedsl_mla_backend import CuteDslMLABackend
return TRTLLMMLABackend(runner, backend="cute-dsl")
return CuteDslMLABackend(runner)
@register_attention_backend("aiter")
@@ -0,0 +1,451 @@
"""
Attention backend for the flashinfer cute-dsl MLA decode kernels with decode
context parallelism (DCP).
Subclasses :class:`TRTLLMMLABackend` (``backend="cute-dsl"``) to reuse its MLA
data preparation, workspace, and prefill plumbing. The flashinfer cute-dsl
monolithic MLA decode kernel natively accepts cyclic DCP metadata
(``enable_dcp`` / ``cp_world`` / ``cp_rank`` / ``causal_seqlens_kv_global``) and
returns the rank-local ``(out, lse)`` needed by the cross-rank merge in
``deepseek_common/attention_forward_methods/forward_mla.py``.
Non-DCP (``dcp_size == 1``) decode falls through to the base cute-dsl path
unchanged. The DCP metadata helpers below are intentionally duplicated from
:mod:`tokenspeed_mla_backend` (they are kernel-agnostic) so that TokenSpeed
stays untouched; both should collapse into the base once the cute-dsl decode
path is stable (see the TODO in tokenspeed_mla_backend.py).
"""
from __future__ import annotations
import logging
import math
from typing import TYPE_CHECKING, Optional
import torch
from sglang.kernels.ops.attention.dcp_kernels import (
create_mla_kv_page_table_for_dcp,
)
from sglang.kernels.ops.attention.fixup_zero_kv import fixup_zero_kv_rows
from sglang.kernels.ops.attention.utils import (
concat_mla_absorb_q_general,
mla_quantize_and_rope_for_fp8,
mla_quantize_without_rope_for_fp8,
)
from sglang.kernels.ops.kvcache.kv_indices import (
get_num_kv_index_blocks_flashmla,
get_num_page_per_block_flashmla,
)
from sglang.srt.environ import envs
from sglang.srt.layers.attention.trtllm_mla_backend import (
TRTLLMMLABackend,
TRTLLMMLAMultiStepDraftBackend,
)
from sglang.srt.layers.dcp.layout import get_dcp_lens
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_flashinfer_available
if is_flashinfer_available():
import flashinfer
if TYPE_CHECKING:
from sglang.srt.layers.radix_attention import RadixAttention
from sglang.srt.model_executor.model_runner import ModelRunner
logger = logging.getLogger(__name__)
# The flashinfer cute-dsl MLA decode kernel returns a natural-log (base-e) LSE,
# whereas sglang's DCP cross-rank merge (forward_mla: dcp_a2a_lse_reduce /
# cp_lse_ag_out_rs_mla) assumes the FlashInfer-MLA/FlashMLA base-2 convention
# (is_lse_base_on_e=False). Multiplying a natural-log LSE by log2(e) rebases it
# to base-2 (the softmax output is base-invariant; only the LSE value changes).
# CONFIRMED base-e (not base-2), so this rebase is required, not optional:
# the flashinfer-dcp-backport public-API unit test asserts the public
# trtllm_batch_decode_with_kv_cache_mla LSE against a torch.logsumexp
# (natural-log) reference at atol=1e-2 and passes (a base-2 LSE would be
# off by 1/ln2 ~= 44%). GPU job 467640:
# tests/attention/test_cute_dsl_mla_dcp*.py 27/27 + 17/17 pass.
_LSE_BASE2_FROM_NATURAL_LOG = math.log2(math.e)
class CuteDslMLABackend(TRTLLMMLABackend):
"""flashinfer cute-dsl MLA decode backend with decode context parallelism."""
def __init__(
self,
model_runner: ModelRunner,
skip_prefill: bool = False,
kv_indptr_buf: Optional[torch.Tensor] = None,
q_indptr_decode_buf: Optional[torch.Tensor] = None,
):
super().__init__(
model_runner,
skip_prefill,
kv_indptr_buf,
q_indptr_decode_buf,
backend="cute-dsl",
)
# ------------------------------------------------------------------
# DCP metadata (rank-local KV lengths + page table).
# Duplicated from TokenspeedMLABackend — kernel-agnostic, keyed only on
# dcp_size / dcp_rank / page_size / req_to_token.
# ------------------------------------------------------------------
def _get_dcp_local_seq_lens(self, seq_lens: torch.Tensor) -> torch.Tensor:
parallel = get_parallel()
if not parallel.dcp_enabled:
return seq_lens
return get_dcp_lens(seq_lens, parallel.dcp_size, parallel.dcp_rank).to(
torch.int32
)
def _get_dcp_local_max_seq_len(self, max_seq_len: int) -> int:
parallel = get_parallel()
if not parallel.dcp_enabled:
return max_seq_len
local_max = max_seq_len // parallel.dcp_size + int(
parallel.dcp_rank < max_seq_len % parallel.dcp_size
)
# A positive scheduling bound is required even when every sequence in a
# padded graph row is empty on this rank.
return max(local_max, 1)
def _fill_dcp_block_kv_indices(
self,
block_kv_indices: torch.Tensor,
req_pool_indices: torch.Tensor,
local_seq_lens: torch.Tensor,
) -> None:
parallel = get_parallel()
pages_per_block = get_num_page_per_block_flashmla(self.page_size)
create_mla_kv_page_table_for_dcp[
(
block_kv_indices.shape[0],
get_num_kv_index_blocks_flashmla(
block_kv_indices.shape[1], self.page_size
),
)
](
self.req_to_token,
req_pool_indices,
local_seq_lens,
block_kv_indices,
self.req_to_token.stride(0),
block_kv_indices.stride(0),
PHYSICAL_PAGE_SIZE=self.page_size,
DCP_SIZE=parallel.dcp_size,
DCP_RANK=parallel.dcp_rank,
PAGES_PER_BLOCK=pages_per_block,
)
def _create_block_kv_indices(
self,
batch_size: int,
max_blocks: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
device: torch.device,
) -> torch.Tensor:
if not get_parallel().dcp_enabled:
return super()._create_block_kv_indices(
batch_size,
max_blocks,
req_pool_indices,
seq_lens,
device,
)
block_kv_indices = torch.full(
(batch_size, max_blocks), -1, dtype=torch.int32, device=device
)
self._fill_dcp_block_kv_indices(
block_kv_indices,
req_pool_indices,
self._get_dcp_local_seq_lens(seq_lens),
)
return block_kv_indices
def _init_cuda_graph_metadata(
self,
bs: int,
num_tokens: int,
forward_mode,
seq_lens: torch.Tensor,
device: torch.device,
):
super()._init_cuda_graph_metadata(
bs, num_tokens, forward_mode, seq_lens, device
)
if get_parallel().dcp_enabled:
if forward_mode.is_target_verify():
self.forward_decode_metadata.global_seq_lens_k = torch.zeros_like(
self.forward_decode_metadata.seq_lens_k
)
self.forward_decode_metadata.max_seq_len_k = (
self._get_dcp_local_max_seq_len(
self.max_context_len
+ (self.num_draft_tokens if forward_mode.is_target_verify() else 0)
)
)
def _apply_cuda_graph_metadata(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
forward_mode,
):
if not get_parallel().dcp_enabled:
return super()._apply_cuda_graph_metadata(
bs,
req_pool_indices,
seq_lens,
forward_mode,
)
metadata = self.decode_cuda_graph_metadata[bs]
if forward_mode.is_target_verify():
torch.add(
seq_lens[:bs],
self.num_draft_tokens,
out=metadata.global_seq_lens_k,
)
metadata.seq_lens_k.copy_(
self._get_dcp_local_seq_lens(metadata.global_seq_lens_k)
)
local_seq_lens = metadata.seq_lens_k
elif forward_mode.is_draft_extend_v2():
num_tokens_per_req = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_req
metadata.sum_seq_lens_q = num_tokens_per_req * bs
seq_lens = seq_lens[:bs]
metadata.seq_lens_k.copy_(seq_lens)
local_seq_lens = self._get_dcp_local_seq_lens(seq_lens)
else:
seq_lens = seq_lens[:bs]
local_seq_lens = self._get_dcp_local_seq_lens(seq_lens)
self._fill_dcp_block_kv_indices(
metadata.block_kv_indices,
req_pool_indices[:bs],
local_seq_lens,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
super().init_forward_metadata(forward_batch)
if (
get_parallel().dcp_enabled
and self.forward_decode_metadata is not None
and (
forward_batch.forward_mode.is_decode_or_idle()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend_v2()
)
):
if forward_batch.forward_mode.is_target_verify():
metadata = self.forward_decode_metadata
metadata.global_seq_lens_k = metadata.seq_lens_k
metadata.seq_lens_k = self._get_dcp_local_seq_lens(
metadata.global_seq_lens_k
)
self.forward_decode_metadata.max_seq_len_k = (
self._get_dcp_local_max_seq_len(
self.forward_decode_metadata.max_seq_len_k
)
)
# ------------------------------------------------------------------
# Kernel + decode forward.
# ------------------------------------------------------------------
def _run_decode_kernel(
self,
query: torch.Tensor,
kv_cache: torch.Tensor,
block_tables: torch.Tensor,
seq_lens: torch.Tensor,
max_seq_len: int,
layer: RadixAttention,
*,
causal_seqs: Optional[torch.Tensor] = None,
cp_world: int = 1,
cp_rank: int = 0,
return_lse: bool = False,
):
"""Call the flashinfer cute-dsl MLA decode kernel.
Without DCP (``cp_world <= 1``) this defers to the base cute-dsl path.
With DCP, ``seq_lens`` are this rank's cyclic-local KV lengths and
``causal_seqs`` the global per-request KV lengths; the kernel returns a
rank-local ``(out, lse)`` (LSE rebased to base-2 for the sglang merge).
"""
if cp_world <= 1:
return super()._run_decode_kernel(
query, kv_cache, block_tables, seq_lens, max_seq_len, layer
)
if causal_seqs is None:
raise ValueError(
"causal_seqs (global per-request KV lengths) is required for DCP "
"MLA decode."
)
bmm1_scale = self._compute_decode_bmm1_scale(layer)
raw_out, lse = flashinfer.decode.trtllm_batch_decode_with_kv_cache_mla(
query=query,
kv_cache=kv_cache,
workspace_buffer=self.workspace_buffer,
qk_nope_head_dim=self.qk_nope_head_dim,
kv_lora_rank=self.kv_lora_rank,
qk_rope_head_dim=self.qk_rope_head_dim,
block_tables=block_tables,
seq_lens=(
seq_lens if seq_lens.dtype == torch.int32 else seq_lens.to(torch.int32)
),
max_seq_len=max_seq_len,
bmm1_scale=bmm1_scale,
skip_softmax_threshold_scale_factor=envs.SGLANG_SKIP_SOFTMAX_DECODE_THRESHOLD_SCALE_FACTOR.get(),
backend="cute-dsl",
enable_dcp=True,
cp_world=cp_world,
cp_rank=cp_rank,
causal_seqlens_kv_global=(
causal_seqs
if causal_seqs.dtype == torch.int32
else causal_seqs.to(torch.int32)
),
return_lse=True, # DCP requires the rank-local LSE for the merge
)
return raw_out, lse * _LSE_BASE2_FROM_NATURAL_LOG
def forward_decode(
self,
q: torch.Tensor, # q_nope
k: torch.Tensor, # k_nope
v: torch.Tensor, # not used in this backend
layer: RadixAttention,
forward_batch: ForwardBatch,
save_kv_cache: bool = True,
q_rope: Optional[torch.Tensor] = None,
k_rope: Optional[torch.Tensor] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
is_neox: Optional[bool] = False,
llama_4_scaling: Optional[torch.Tensor] = None,
):
parallel = get_parallel()
if not parallel.dcp_enabled:
return super().forward_decode(
q,
k,
v,
layer,
forward_batch,
save_kv_cache,
q_rope,
k_rope,
cos_sin_cache,
is_neox,
llama_4_scaling,
)
# Query / KV preparation mirrors the base cute-dsl decode (both FP16 and
# FP8 KV), then swaps to the DCP kernel call + rank-local return.
merge_query = q_rope is not None
if self.data_type == torch.float8_e4m3fn:
assert q_rope is not None and k_rope is not None
if cos_sin_cache is None:
q, k, k_rope = mla_quantize_without_rope_for_fp8(
q, q_rope, k.squeeze(1), k_rope.squeeze(1)
)
else:
q, k, k_rope = mla_quantize_and_rope_for_fp8(
q,
q_rope,
k.squeeze(1),
k_rope.squeeze(1),
forward_batch.positions,
cos_sin_cache,
is_neox,
self.kv_lora_rank,
self.qk_rope_head_dim,
)
merge_query = False
if save_kv_cache:
assert k is not None and k_rope is not None
self.token_to_kv_pool.set_mla_kv_buffer(
layer, forward_batch.out_cache_loc, k, k_rope
)
if merge_query:
q_nope = q.view(-1, layer.tp_q_head_num, layer.v_head_dim)
q_rope_reshaped = q_rope.view(
-1, layer.tp_q_head_num, layer.head_dim - layer.v_head_dim
)
query = concat_mla_absorb_q_general(q_nope, q_rope_reshaped)
else:
query = q.view(-1, layer.tp_q_head_num, layer.head_dim)
if llama_4_scaling is not None:
query = (query.to(self.q_data_type) * llama_4_scaling).to(self.data_type)
if query.dim() == 3:
query = query.unsqueeze(1)
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
kv_cache = k_cache.view(-1, self.page_size, self.kv_cache_dim).unsqueeze(1)
metadata = (
getattr(forward_batch, "decode_trtllm_mla_metadata", None)
or self.forward_decode_metadata
)
metadata_batch_size = getattr(metadata, "batch_size", None)
if (
metadata_batch_size is not None
and metadata_batch_size < forward_batch.batch_size
):
self.init_forward_metadata(forward_batch)
metadata = forward_batch.decode_trtllm_mla_metadata
global_seq_lens = forward_batch.seq_lens[: forward_batch.batch_size]
local_seq_lens = self._get_dcp_local_seq_lens(global_seq_lens)
raw_out, lse = self._run_decode_kernel(
query=query,
kv_cache=kv_cache,
block_tables=metadata.block_kv_indices,
seq_lens=local_seq_lens,
max_seq_len=metadata.max_seq_len_k,
layer=layer,
causal_seqs=global_seq_lens,
cp_world=parallel.dcp_size,
cp_rank=parallel.dcp_rank,
return_lse=True,
)
output = raw_out.view(-1, layer.tp_q_head_num, layer.v_head_dim)
lse = lse.view(-1, layer.tp_q_head_num)
# Zero-KV rows (a request this rank owns no cyclic slice for) get a
# neutral (out=0, lse=-inf) state so the cross-rank merge ignores them.
fixup_zero_kv_rows(
output,
lse,
local_seq_lens,
self.q_indptr_decode[: forward_batch.batch_size + 1],
1,
)
return output.flatten(1), lse
class CuteDslMLAMultiStepDraftBackend(TRTLLMMLAMultiStepDraftBackend):
"""Multi-step draft backend for cutedsl_mla used by EAGLE / DSPARK."""
def __init__(
self, model_runner: ModelRunner, topk: int, speculative_num_steps: int
):
super().__init__(model_runner, topk, speculative_num_steps)
# Parent populates self.attn_backends with TRT-LLM instances; replace
# them with cute-dsl instances sharing the parent's index buffers.
for i in range(self.speculative_num_steps - 1):
self.attn_backends[i] = CuteDslMLABackend(
model_runner,
skip_prefill=True,
kv_indptr_buf=self.kv_indptr[i],
q_indptr_decode_buf=self.q_indptr_decode,
)
@@ -49,11 +49,11 @@ from sglang.srt.speculative.spec_utils import (
generate_draft_decode_kv_indices,
)
from sglang.srt.utils import (
get_cuda_graph_max_batch_size,
get_int_env_var,
is_flashinfer_available,
is_sm100_supported,
next_power_of_2,
require_gathered_buffer,
)
if TYPE_CHECKING:
@@ -63,18 +63,6 @@ if TYPE_CHECKING:
logger = logging.getLogger(__name__)
def _cuda_graph_capture_max_bs(server_args, max_bs: int) -> int:
"""Pad max_bs to the alignment cuda-graph capture uses (see get_batch_sizes_to_capture)."""
mul_base = 1
if server_args.enable_two_batch_overlap:
mul_base *= 2
if require_gathered_buffer(server_args):
mul_base *= get_parallel().attn_tp_size
if mul_base % get_parallel().attn_cp_size != 0:
mul_base *= get_parallel().attn_cp_size
return (max_bs + mul_base - 1) // mul_base * mul_base
if envs.SGLANG_ENABLE_TORCH_COMPILE.get():
torch._logging.set_logs(dynamo=logging.ERROR)
torch._dynamo.config.suppress_errors = True
@@ -447,7 +435,7 @@ class FlashInferAttnBackend(AttentionBackend):
)
else:
self.workspace_buffer = global_workspace_buffer
max_bs = _cuda_graph_capture_max_bs(
max_bs = get_cuda_graph_max_batch_size(
model_runner.server_args, model_runner.req_to_token_pool.size
)
if kv_indptr_buf is None:
@@ -2230,7 +2218,7 @@ class FlashInferMultiStepDraftBackend:
self.generate_draft_decode_kv_indices = generate_draft_decode_kv_indices
self.page_size = model_runner.page_size
max_bs = _cuda_graph_capture_max_bs(
max_bs = get_cuda_graph_max_batch_size(
model_runner.server_args, model_runner.req_to_token_pool.size * self.topk
)
self.kv_indptr = torch.zeros(
@@ -22,10 +22,10 @@ from __future__ import annotations
"""Attention backend for the tokenspeed-mla CuTe DSL kernels on Blackwell.
Subclasses :class:`TRTLLMMLABackend` and overrides only ``_run_decode_kernel``
and ``_run_prefill_kernel``. All metadata, KV-cache layout, CUDA-graph
plumbing, FP8 quantize/rope, draft-extend padding, and chunked-prefix
dispatch are inherited unchanged from the parent.
Subclasses :class:`TRTLLMMLABackend` to share its MLA data preparation and
prefill plumbing. Decode-context parallelism is implemented here because the
TokenSpeed decode kernel natively accepts CP rank/world metadata and returns
the partial log-sum-exp needed by the cross-rank merge.
"""
import logging
@@ -34,14 +34,28 @@ from typing import TYPE_CHECKING, Optional
import torch
from sglang.kernels.jit.utils import is_arch_support_pdl
from sglang.kernels.ops.attention.dcp_kernels import (
create_mla_kv_page_table_for_dcp,
)
from sglang.kernels.ops.attention.fixup_zero_kv import fixup_zero_kv_rows
from sglang.kernels.ops.attention.mla_kv_pack_quantize_fp8 import (
mla_kv_pack_quantize_fp8,
)
from sglang.kernels.ops.attention.utils import (
mla_quantize_and_rope_for_fp8,
mla_quantize_without_rope_for_fp8,
)
from sglang.kernels.ops.kvcache.kv_indices import (
get_num_kv_index_blocks_flashmla,
get_num_page_per_block_flashmla,
)
from sglang.kernels.ops.quantization.fp8_quantize import fp8_quantize
from sglang.srt.layers.attention.trtllm_mla_backend import (
TRTLLMMLABackend,
TRTLLMMLAMultiStepDraftBackend,
)
from sglang.srt.layers.dcp.layout import get_dcp_lens
from sglang.srt.runtime_context import get_parallel
from sglang.srt.utils import is_flashinfer_available, is_tokenspeed_mla_available
if is_flashinfer_available():
@@ -273,6 +287,162 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
k_nope, k_pe, v, enable_pdl=is_arch_support_pdl()
)
def _get_dcp_local_seq_lens(self, seq_lens: torch.Tensor) -> torch.Tensor:
parallel = get_parallel()
if not parallel.dcp_enabled:
return seq_lens
return get_dcp_lens(seq_lens, parallel.dcp_size, parallel.dcp_rank).to(
torch.int32
)
def _get_dcp_local_max_seq_len(self, max_seq_len: int) -> int:
parallel = get_parallel()
if not parallel.dcp_enabled:
return max_seq_len
local_max = max_seq_len // parallel.dcp_size + int(
parallel.dcp_rank < max_seq_len % parallel.dcp_size
)
# TokenSpeed requires a positive scheduling bound even when every
# sequence in a padded graph row is empty on this rank.
return max(local_max, 1)
def _fill_dcp_block_kv_indices(
self,
block_kv_indices: torch.Tensor,
req_pool_indices: torch.Tensor,
local_seq_lens: torch.Tensor,
) -> None:
parallel = get_parallel()
pages_per_block = get_num_page_per_block_flashmla(self.page_size)
create_mla_kv_page_table_for_dcp[
(
block_kv_indices.shape[0],
get_num_kv_index_blocks_flashmla(
block_kv_indices.shape[1], self.page_size
),
)
](
self.req_to_token,
req_pool_indices,
local_seq_lens,
block_kv_indices,
self.req_to_token.stride(0),
block_kv_indices.stride(0),
PHYSICAL_PAGE_SIZE=self.page_size,
DCP_SIZE=parallel.dcp_size,
DCP_RANK=parallel.dcp_rank,
PAGES_PER_BLOCK=pages_per_block,
)
def _create_block_kv_indices(
self,
batch_size: int,
max_blocks: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
device: torch.device,
) -> torch.Tensor:
if not get_parallel().dcp_enabled:
return super()._create_block_kv_indices(
batch_size,
max_blocks,
req_pool_indices,
seq_lens,
device,
)
block_kv_indices = torch.full(
(batch_size, max_blocks), -1, dtype=torch.int32, device=device
)
self._fill_dcp_block_kv_indices(
block_kv_indices,
req_pool_indices,
self._get_dcp_local_seq_lens(seq_lens),
)
return block_kv_indices
def _init_cuda_graph_metadata(
self,
bs: int,
num_tokens: int,
forward_mode,
seq_lens: torch.Tensor,
device: torch.device,
):
super()._init_cuda_graph_metadata(
bs, num_tokens, forward_mode, seq_lens, device
)
if get_parallel().dcp_enabled:
self.forward_decode_metadata.max_seq_len_k = (
self._get_dcp_local_max_seq_len(
self.max_context_len
+ (self.num_draft_tokens if forward_mode.is_target_verify() else 0)
)
)
def _apply_cuda_graph_metadata(
self,
bs: int,
req_pool_indices: torch.Tensor,
seq_lens: torch.Tensor,
forward_mode,
):
if not get_parallel().dcp_enabled:
return super()._apply_cuda_graph_metadata(
bs, req_pool_indices, seq_lens, forward_mode
)
metadata = self.decode_cuda_graph_metadata[bs]
if forward_mode.is_target_verify():
torch.add(
seq_lens[:bs],
self.num_draft_tokens,
out=metadata.global_seq_lens_k,
)
metadata.seq_lens_k.copy_(
self._get_dcp_local_seq_lens(metadata.global_seq_lens_k)
)
local_seq_lens = metadata.seq_lens_k
elif forward_mode.is_draft_extend_v2():
num_tokens_per_req = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_req
metadata.sum_seq_lens_q = num_tokens_per_req * bs
seq_lens = seq_lens[:bs]
metadata.seq_lens_k.copy_(seq_lens)
local_seq_lens = self._get_dcp_local_seq_lens(seq_lens)
else:
seq_lens = seq_lens[:bs]
local_seq_lens = self._get_dcp_local_seq_lens(seq_lens)
self._fill_dcp_block_kv_indices(
metadata.block_kv_indices,
req_pool_indices[:bs],
local_seq_lens,
)
def init_forward_metadata(self, forward_batch: ForwardBatch):
super().init_forward_metadata(forward_batch)
if (
get_parallel().dcp_enabled
and self.forward_decode_metadata is not None
and (
forward_batch.forward_mode.is_decode_or_idle()
or forward_batch.forward_mode.is_target_verify()
or forward_batch.forward_mode.is_draft_extend_v2()
)
):
if forward_batch.forward_mode.is_target_verify():
metadata = self.forward_decode_metadata
metadata.global_seq_lens_k = metadata.seq_lens_k
metadata.seq_lens_k = self._get_dcp_local_seq_lens(
metadata.global_seq_lens_k
)
self.forward_decode_metadata.max_seq_len_k = (
self._get_dcp_local_max_seq_len(
self.forward_decode_metadata.max_seq_len_k
)
)
def _run_decode_kernel(
self,
query: torch.Tensor,
@@ -281,7 +451,12 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
seq_lens: torch.Tensor,
max_seq_len: int,
layer: RadixAttention,
) -> torch.Tensor:
*,
causal_seqs: Optional[torch.Tensor] = None,
cp_world: int = 1,
cp_rank: int = 0,
return_lse: bool = False,
):
k_scale = getattr(layer, "k_scale_float", None)
if k_scale is None:
k_scale = 1.0
@@ -303,8 +478,111 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
softmax_scale=softmax_scale,
output_scale=output_scale,
enable_pdl=is_arch_support_pdl(),
return_lse=return_lse,
causal_seqs=causal_seqs,
cp_world=cp_world,
cp_rank=cp_rank,
)
def forward_decode(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
layer: RadixAttention,
forward_batch: ForwardBatch,
save_kv_cache: bool = True,
q_rope: Optional[torch.Tensor] = None,
k_rope: Optional[torch.Tensor] = None,
cos_sin_cache: Optional[torch.Tensor] = None,
is_neox: Optional[bool] = False,
llama_4_scaling: Optional[torch.Tensor] = None,
):
parallel = get_parallel()
if not parallel.dcp_enabled:
return super().forward_decode(
q,
k,
v,
layer,
forward_batch,
save_kv_cache,
q_rope,
k_rope,
cos_sin_cache,
is_neox,
llama_4_scaling,
)
assert q_rope is not None and k_rope is not None
if cos_sin_cache is None:
q, k, k_rope = mla_quantize_without_rope_for_fp8(
q, q_rope, k.squeeze(1), k_rope.squeeze(1)
)
else:
q, k, k_rope = mla_quantize_and_rope_for_fp8(
q,
q_rope,
k.squeeze(1),
k_rope.squeeze(1),
forward_batch.positions,
cos_sin_cache,
is_neox,
self.kv_lora_rank,
self.qk_rope_head_dim,
)
if save_kv_cache:
self.token_to_kv_pool.set_mla_kv_buffer(
layer, forward_batch.out_cache_loc, k, k_rope
)
query = q.view(-1, layer.tp_q_head_num, layer.head_dim)
if llama_4_scaling is not None:
query = (query.to(self.q_data_type) * llama_4_scaling).to(self.data_type)
if query.dim() == 3:
query = query.unsqueeze(1)
k_cache = self.token_to_kv_pool.get_key_buffer(layer.layer_id)
kv_cache = k_cache.view(-1, self.page_size, self.kv_cache_dim).unsqueeze(1)
metadata = (
getattr(forward_batch, "decode_trtllm_mla_metadata", None)
or self.forward_decode_metadata
)
metadata_batch_size = getattr(metadata, "batch_size", None)
if (
metadata_batch_size is not None
and metadata_batch_size < forward_batch.batch_size
):
self.init_forward_metadata(forward_batch)
metadata = forward_batch.decode_trtllm_mla_metadata
global_seq_lens = forward_batch.seq_lens[: forward_batch.batch_size]
local_seq_lens = self._get_dcp_local_seq_lens(global_seq_lens)
raw_out, lse = self._run_decode_kernel(
query=query,
kv_cache=kv_cache,
block_tables=metadata.block_kv_indices,
seq_lens=local_seq_lens,
max_seq_len=metadata.max_seq_len_k,
layer=layer,
causal_seqs=global_seq_lens,
cp_world=parallel.dcp_size,
cp_rank=parallel.dcp_rank,
return_lse=True,
)
output = raw_out.view(-1, layer.tp_q_head_num, layer.v_head_dim)
lse = lse.view(-1, layer.tp_q_head_num)
fixup_zero_kv_rows(
output,
lse,
local_seq_lens,
self.q_indptr_decode[: forward_batch.batch_size + 1],
1,
)
return output.flatten(1), lse
def _run_prefill_kernel(
self,
q: torch.Tensor,
@@ -323,6 +601,14 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
o_sf_scale: float = 1.0,
): # Q/K/V arrive already in FP8 via the model-side fused path
# (prepare_prefill_qkv / pack_prefix_chunk_kv); no quantize here.
# Hybrid MLA models resolve the model-side hook through the outer
# HybridLinearAttnBackend, so their fallback MHA path can pass V as a
# last-dimension slice of kv_b_proj (stride(-2) > size(-1)). The
# TokenSpeed prefill kernel requires dense Q/K/V layouts even though
# the public wrapper accepts arbitrary torch tensors.
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
return tokenspeed_mla.tokenspeed_mla_prefill(
query=q,
key=k,
@@ -150,6 +150,7 @@ class TRTLLMMLADecodeMetadata:
cu_seqlens_q: Optional[torch.Tensor] = None
seq_lens_q: Optional[torch.Tensor] = None
seq_lens_k: Optional[torch.Tensor] = None
global_seq_lens_k: Optional[torch.Tensor] = None
class TRTLLMMLABackend(FlashInferMLAAttnBackend):
@@ -382,6 +383,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if forward_mode.is_target_verify():
metadata.seq_lens_k = torch.zeros((bs,), dtype=torch.int32, device=device)
metadata.global_seq_lens_k = torch.zeros(
(bs,), dtype=torch.int32, device=device
)
elif forward_mode.is_draft_extend_v2():
num_tokens_per_req = self.num_draft_tokens
metadata.max_seq_len_q = num_tokens_per_req
@@ -423,7 +427,12 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
if forward_mode.is_target_verify():
# Intentional int64 -> int32 same-kind out= downcast.
torch.add(seq_lens[:bs], self.num_draft_tokens, out=metadata.seq_lens_k)
torch.add(
seq_lens[:bs],
self.num_draft_tokens,
out=metadata.global_seq_lens_k,
)
metadata.seq_lens_k.copy_(metadata.global_seq_lens_k)
seq_lens = metadata.seq_lens_k
elif forward_mode.is_draft_extend_v2():
num_tokens_per_req = self.num_draft_tokens
@@ -569,6 +578,9 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
max_seq = max_seq + self.num_draft_tokens
seq_lens = seq_lens + self.num_draft_tokens
self.forward_decode_metadata.seq_lens_k = seq_lens.to(torch.int32)
self.forward_decode_metadata.global_seq_lens_k = (
self.forward_decode_metadata.seq_lens_k
)
elif forward_batch.forward_mode.is_draft_extend_v2():
sum_seq_lens_q = sum(forward_batch.extend_seq_lens_cpu)
max_seq_len_q = max(forward_batch.extend_seq_lens_cpu)
@@ -951,8 +963,10 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
q = q.to(self.data_type)
if forward_batch.forward_mode.is_target_verify():
max_seq_len = (
metadata.max_seq_len_k + forward_batch.spec_info.draft_token_num
draft_token_num = forward_batch.spec_info.draft_token_num
dcp_enabled = get_parallel().dcp_enabled
max_seq_len = metadata.max_seq_len_k + (
0 if dcp_enabled else draft_token_num
)
# For target_verify, all sequences have the same number of draft tokens
q = q.view(bs, -1, layer.tp_q_head_num, layer.head_dim)
@@ -1006,6 +1020,44 @@ class TRTLLMMLABackend(FlashInferMLAAttnBackend):
assert kv_cache.dtype == self.data_type
if (
forward_batch.forward_mode.is_target_verify()
and get_parallel().dcp_enabled
):
raw_out, lse = self._run_decode_kernel(
query=q,
kv_cache=kv_cache,
block_tables=metadata.block_kv_indices,
seq_lens=metadata.seq_lens_k,
max_seq_len=max_seq_len,
layer=layer,
causal_seqs=metadata.global_seq_lens_k,
cp_world=get_parallel().dcp_size,
cp_rank=get_parallel().dcp_rank,
return_lse=True,
)
output = raw_out.view(
bs * draft_token_num,
layer.tp_q_head_num,
layer.v_head_dim,
)
lse = lse.view(bs * draft_token_num, layer.tp_q_head_num)
dense_q_indptr = torch.arange(
0,
(bs + 1) * draft_token_num,
draft_token_num,
dtype=torch.int32,
device=q.device,
)
fixup_zero_kv_rows(
output,
lse,
metadata.seq_lens_k,
dense_q_indptr,
draft_token_num,
)
return output.flatten(1), lse
raw_out = self._run_decode_kernel(
query=q,
kv_cache=kv_cache,
@@ -3634,6 +3634,12 @@ class HybridLinearKVPool(KVCache):
def get_kv_size_bytes(self):
return self.full_kv_pool.get_kv_size_bytes()
def get_kv_buffer_shape(self) -> Tuple[torch.Size, torch.Size]:
# Hybrid layer ids are global model-layer ids, while the backing pool
# is dense over only full-attention layers. Shape discovery does not
# need a global layer lookup, so delegate it to that backing pool.
return self.full_kv_pool.get_kv_buffer_shape()
def get_contiguous_buf_infos(self):
return self.full_kv_pool.get_contiguous_buf_infos()
@@ -23,8 +23,11 @@ from contextlib import contextmanager
from typing import TYPE_CHECKING, Any, List, Sequence, Tuple
from sglang.srt.model_executor.runner.base_runner import BaseRunner
from sglang.srt.runtime_context import get_flags, get_parallel
from sglang.srt.utils import require_gathered_buffer
from sglang.srt.runtime_context import get_flags
from sglang.srt.utils import (
get_cuda_graph_batch_size_alignment,
get_cuda_graph_max_batch_size,
)
if TYPE_CHECKING:
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
@@ -68,22 +71,15 @@ def get_batch_sizes_to_capture(
capture_bs = list(server_args.cuda_graph_config.decode.bs)
num_max_requests = model_runner.req_to_token_pool.size
mul_base = 1
mul_base = get_cuda_graph_batch_size_alignment(server_args)
# TBO splits each request's rows across two micro-batches, so the
# alignment constraint applies per request rather than per token row.
alignment_width = captured_req_width
if server_args.enable_two_batch_overlap:
mul_base *= 2
alignment_width = 1
if require_gathered_buffer(server_args):
mul_base *= get_parallel().attn_tp_size
if mul_base % get_parallel().attn_cp_size != 0:
mul_base *= get_parallel().attn_cp_size
# pad `num_max_requests` to avoid being filtered out
num_max_requests = (num_max_requests + mul_base - 1) // mul_base * mul_base
num_max_requests = get_cuda_graph_max_batch_size(server_args, num_max_requests)
if max(capture_bs) > num_max_requests:
# In some cases (e.g., with a small GPU or --max-running-requests), the #max-running-requests
# is very small. We add more values here to make sure we capture the maximum bs.
@@ -50,7 +50,11 @@ from sglang.srt.model_executor.runner_backend_utils.tc_piecewise_cuda_graph impo
set_tc_piecewise_forward_context,
)
from sglang.srt.utils import is_hip
from sglang.srt.utils.common import ceil_align, require_mlp_sync
from sglang.srt.utils.common import (
ceil_align,
get_eager_max_batch_size,
require_mlp_sync,
)
logger = logging.getLogger(__name__)
@@ -101,14 +105,12 @@ class EagerRunner(BaseRunner):
# (expand_for_topk_draft) before the eager fallback.
max_bs *= sa.speculative_eagle_topk
# Mirror prepare_mlp_sync_batch padding so the registry holds what load_batch copies.
if require_mlp_sync(sa):
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
max_bs = ceil_align(max_bs, self.attn_tp_size)
max_bs = ceil_align(max_bs, get_cp_padding_align_size())
max_bs = get_eager_max_batch_size(sa, max_bs)
prefill_ceiling = max(mr.max_total_num_tokens, sa.max_prefill_buffer_tokens())
max_num_token = max(prefill_ceiling, max_bs * num_tokens_per_req)
if require_mlp_sync(sa):
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
max_num_token = ceil_align(max_num_token, self.attn_tp_size)
max_num_token = ceil_align(max_num_token, get_cp_padding_align_size())
self._eager_max_bs = max_bs
@@ -261,7 +263,9 @@ class EagerRunner(BaseRunner):
forward_batch = self.load_batch(forward_batch, pp_proxy_tensors)
if forward_batch.needs_forward_metadata_init():
if hasattr(model_runner.model, "prepare_context_parallel_metadata_for_dcp"):
if model_runner.dcp_size > 1 and hasattr(
model_runner.model, "prepare_context_parallel_metadata_for_dcp"
):
# prepare kv cache buffer for dcp to gather kv cache
forward_batch.attn_dcp_metadata = (
model_runner.model.prepare_context_parallel_metadata_for_dcp(
@@ -88,6 +88,25 @@ class MlaBmmFusionPlan:
attn_output_buf: torch.Tensor
def _is_dcp_mla_decode_phase(forward_batch: ForwardBatch) -> bool:
if not get_parallel().dcp_enabled:
return False
if forward_batch.forward_mode.is_decode():
return True
if not forward_batch.forward_mode.is_target_verify() or not _is_cuda:
return False
server_args = get_server_args()
decode_backend = (
server_args.decode_attention_backend or server_args.attention_backend
)
return (
server_args.speculative_algorithm == "DSPARK"
and server_args.speculative_attention_mode == "decode"
and decode_backend in ("tokenspeed_mla", "cutedsl_mla")
)
if _is_cuda:
from sglang.kernels.ops.gemm import bmm_fp8
@@ -254,8 +273,7 @@ class DeepseekMLAForwardMixin:
# weights and skip the per-layer Q all-gather (bf16 decode absorb only).
q_replicate_active = (
get_server_args().dcp_replicate_q_proj
and get_parallel().dcp_enabled
and forward_batch.forward_mode.is_decode()
and _is_dcp_mla_decode_phase(forward_batch)
and not self.use_deep_gemm_bmm
and self.w_kc_qrep is not None
and self.q_b_proj_qrep_weight is not None
@@ -595,12 +613,12 @@ class DeepseekMLAForwardMixin:
# all_gather q_pe, q_nope_out,take tp8 as an example q_pe [B, H, ROPE_DIM], q_nope_out [B, H, NOPE_DIM] gathered to [B, H * dcp_world_size, ROPE_DIM] [B, H * dcp_world_size, NOPE_DIM] for decode batch, and all gather k_pe, k_nope for extend batch.
if get_parallel().dcp_enabled:
if forward_batch.forward_mode.is_decode() and not q_replicate_active:
# if forward_batch.forward_mode is decode, gather q
q_nope_out, q_pe = all_gather_q_for_mla_decode(
q_nope_out=q_nope_out,
q_pe=q_pe,
)
if _is_dcp_mla_decode_phase(forward_batch):
if not q_replicate_active:
q_nope_out, q_pe = all_gather_q_for_mla_decode(
q_nope_out=q_nope_out,
q_pe=q_pe,
)
elif forward_batch.forward_mode.is_extend():
# for extend, gather kv
all_gather_kv_cache_for_mla_extend(
@@ -748,10 +766,7 @@ class DeepseekMLAForwardMixin:
topk_indices=topk_indices,
)
attn_output = fusion_plan.attn_output_buf
elif (
forward_batch.forward_mode.is_decode()
and get_parallel().dcp_enabled
):
elif _is_dcp_mla_decode_phase(forward_batch):
# set return_lse=True to correct attn_output
attn_output, lse = self.attn_mqa_for_dcp_decode(
q_nope_out,
@@ -825,7 +840,7 @@ class DeepseekMLAForwardMixin:
)
# correct attn_output with respect to lse from other ranks
if forward_batch.forward_mode.is_decode() and get_parallel().dcp_enabled:
if _is_dcp_mla_decode_phase(forward_batch):
attn_output = attn_output.view(
-1,
self.num_local_heads * get_parallel().attn_dcp_size,
+82 -11
View File
@@ -16,6 +16,7 @@ from sglang.srt.distributed import (
tensor_model_parallel_all_reduce,
)
from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder
from sglang.srt.layers.dcp.planner import prepare_decode_context_parallel_metadata
from sglang.srt.layers.layernorm import RMSNorm
from sglang.srt.layers.linear import (
ColumnParallelBatchedLinear,
@@ -52,6 +53,20 @@ from sglang.srt.utils import make_layers
from sglang.srt.utils.common import BumpAllocator, add_prefix, set_weight_attrs
def _get_kda_local_num_heads(num_heads: int, tp_size: int) -> int:
if num_heads % tp_size != 0:
raise ValueError(
f"KDA num_heads ({num_heads}) must be divisible by global tp_size ({tp_size})"
)
return num_heads // tp_size
def _materialize_residual_stream(
hidden_states: torch.Tensor, residual: Optional[torch.Tensor]
) -> torch.Tensor:
return hidden_states if residual is None else hidden_states + residual
class KimiMoE(nn.Module):
def __init__(
self,
@@ -188,8 +203,7 @@ class KimiDeltaAttention(nn.Module):
self.head_v_dim = config.linear_attn_config["head_dim"]
self.layer_idx = layer_idx
self.prefix = prefix
assert self.num_heads % self.tp_size == 0
self.local_num_heads = divide(self.num_heads, self.tp_size)
self.local_num_heads = _get_kda_local_num_heads(self.num_heads, self.tp_size)
projection_size = self.head_dim * self.num_heads
self.conv_size = config.linear_attn_config["short_conv_kernel_size"]
@@ -317,9 +331,9 @@ class KimiDeltaAttention(nn.Module):
self.attn = RadixLinearAttention(
layer_id=self.layer_idx,
num_q_heads=self.num_k_heads // self.attn_tp_size,
num_k_heads=self.num_k_heads // self.attn_tp_size,
num_v_heads=self.num_v_heads // self.attn_tp_size,
num_q_heads=_get_kda_local_num_heads(self.num_k_heads, self.tp_size),
num_k_heads=_get_kda_local_num_heads(self.num_k_heads, self.tp_size),
num_v_heads=_get_kda_local_num_heads(self.num_v_heads, self.tp_size),
head_q_dim=self.head_k_dim,
head_k_dim=self.head_k_dim,
head_v_dim=self.head_v_dim,
@@ -519,6 +533,7 @@ class KimiLinearModel(nn.Module):
self.padding_idx = config.pad_token_id
self.vocab_size = config.vocab_size
self.pp_group = get_pp_group()
self.dspark_layers_to_capture: Optional[list[int]] = None
if self.pp_group.is_first_rank:
self.embed_tokens = VocabParallelEmbedding(
@@ -581,7 +596,6 @@ class KimiLinearModel(nn.Module):
dtype=torch.float32,
device=device,
)
# TODO: capture aux hidden states
aux_hidden_states = []
for i in range(self.start_layer, self.end_layer):
ctx = get_global_expert_distribution_recorder().with_current_layer(i)
@@ -594,6 +608,13 @@ class KimiLinearModel(nn.Module):
residual=residual,
zero_allocator=zero_allocator,
)
if (
self.dspark_layers_to_capture is not None
and i in self.dspark_layers_to_capture
):
aux_hidden_states.append(
_materialize_residual_stream(hidden_states, residual)
)
if not self.pp_group.is_last_rank:
return PPProxyTensors(
@@ -609,10 +630,9 @@ class KimiLinearModel(nn.Module):
else:
hidden_states, _ = self.norm(hidden_states, residual)
if len(aux_hidden_states) == 0:
return hidden_states
return hidden_states, aux_hidden_states
if self.dspark_layers_to_capture is not None:
return hidden_states, aux_hidden_states
return hidden_states
class KimiLinearForCausalLM(nn.Module):
@@ -642,6 +662,22 @@ class KimiLinearForCausalLM(nn.Module):
self.lm_head = PPMissingLayer()
logit_scale = getattr(self.config, "logit_scale", 1.0)
self.logits_processor = LogitsProcessor(config=config, logit_scale=logit_scale)
self.capture_aux_hidden_states = False
def get_input_embeddings(self):
return self.model.embed_tokens
def set_dspark_layers_to_capture(self, layer_ids: list[int]) -> None:
if self.pp_group.world_size > 1:
raise NotImplementedError("DSPARK aux hidden capture requires PP=1.")
if not self.pp_group.is_last_rank:
return
if layer_ids is None:
raise ValueError(
"DSPARK requires explicit layer_ids for aux hidden capture."
)
self.capture_aux_hidden_states = True
self.model.dspark_layers_to_capture = list(layer_ids)
@torch.no_grad()
def forward(
@@ -660,12 +696,47 @@ class KimiLinearForCausalLM(nn.Module):
pp_proxy_tensors,
)
if self.pp_group.is_last_rank:
aux_hidden_states = None
if self.capture_aux_hidden_states:
hidden_states, aux_hidden_states = hidden_states
return self.logits_processor(
input_ids, hidden_states, self.lm_head, forward_batch
input_ids,
hidden_states,
self.lm_head,
forward_batch,
aux_hidden_states,
)
else:
return hidden_states
def prepare_context_parallel_metadata_for_dcp(
self,
seq_lens: torch.Tensor,
extend_prefix_lens: torch.Tensor,
extend_prefix_lens_cpu: torch.Tensor,
extend_seq_lens: torch.Tensor,
req_pool_indices: torch.Tensor,
req_to_token: torch.Tensor,
seq_lens_sum: int,
kv_buffer_shape: torch.Size,
kv_cache_dtype,
kv_cache_device,
create_chunked_prefix_cache_kv_indices_fn,
):
return prepare_decode_context_parallel_metadata(
seq_lens=seq_lens,
extend_prefix_lens=extend_prefix_lens,
extend_prefix_lens_cpu=extend_prefix_lens_cpu,
extend_seq_lens=extend_seq_lens,
req_pool_indices=req_pool_indices,
req_to_token=req_to_token,
seq_lens_sum=seq_lens_sum,
kv_buffer_shape=kv_buffer_shape,
kv_cache_dtype=kv_cache_dtype,
kv_cache_device=kv_cache_device,
create_chunked_prefix_cache_kv_indices_fn=create_chunked_prefix_cache_kv_indices_fn,
)
def _is_non_local_pp_weight(self, name: str) -> bool:
if self.pp_group.world_size == 1:
return False
+35 -17
View File
@@ -1038,14 +1038,6 @@ class ServerArgs:
),
NS("parallel"),
] = 1
dcp_size: A[
int,
Arg(
help="The decode context parallelism size.",
aliases=["--decode-context-parallel-size"],
),
NS("parallel"),
] = 1
dwdp_size: A[
int,
Arg(
@@ -1064,20 +1056,24 @@ class ServerArgs:
"combine), or 'fi_a2a' (FlashInfer MNNVL All-to-All kernel; requires "
"SM90+ and MNNVL fabric memory, e.g. GB200 NVL72).",
choices=["ag_rs", "a2a", "fi_a2a"],
resolvable=True,
),
NS("parallel"),
] = "ag_rs"
dcp_replicate_q_proj: A[
bool,
Optional[bool],
Arg(
help="For MLA decode context parallelism with the a2a/fi_a2a "
"backend: replicate the Q projection so each DCP rank computes the "
"full-head query locally (redundant projection compute), eliminating "
"the per-layer head-dim all-gather of Q. Trades a small amount of "
"extra GEMM for one fewer collective per layer.",
"extra GEMM for one fewer collective per layer. Use "
"--no-dcp-replicate-q-proj to disable the model-specific default.",
action=argparse.BooleanOptionalAction,
resolvable=True,
),
NS("parallel"),
] = False
] = None
enable_prefill_cp: A[
bool,
"Enable context parallelism for the prefill phase. Select the layout with --cp-strategy.",
@@ -3763,13 +3759,35 @@ class ServerArgs:
return
elif is_cuda():
if self.speculative_algorithm is not None:
raise ValueError(
"Decode context parallel (--dcp-size / "
"--decode-context-parallel-size > 1) on CUDA platform "
"does not support any speculative algorithm, but got "
f"dcp_size={self.dcp_size} on a CUDA platform with "
"speculative decoding enabled."
model_arches = self.get_model_config().hf_config.architectures
decode_backend = self.decode_attention_backend or self.attention_backend
kimi_linear_dspark = (
self.speculative_algorithm == "DSPARK"
and "KimiLinearForCausalLM" in model_arches
and self.speculative_attention_mode == "decode"
and decode_backend in ("tokenspeed_mla", "cutedsl_mla")
)
if kimi_linear_dspark:
ragged_verify_mode = envs.SGLANG_RAGGED_VERIFY_MODE.get()
if ragged_verify_mode != "static":
raise ValueError(
"Kimi Linear DCP + DSPARK currently requires "
"SGLANG_RAGGED_VERIFY_MODE=static, but got "
f"{ragged_verify_mode!r}."
)
else:
raise ValueError(
"Decode context parallel (--dcp-size / "
"--decode-context-parallel-size > 1) with speculative "
"decoding on CUDA is supported only for Kimi Linear + "
"DSPARK + --speculative-attention-mode decode + "
"tokenspeed_mla, or experimental cutedsl_mla, but got "
f"architectures={model_arches}, "
f"speculative_algorithm={self.speculative_algorithm!r}, "
"speculative_attention_mode="
f"{self.speculative_attention_mode!r}, "
f"decode_attention_backend={decode_backend!r}."
)
else:
raise ValueError(
"Decode context parallel (--dcp-size / "
+12 -1
View File
@@ -239,7 +239,18 @@ class DraftBackendFactory:
)
def _create_cutedsl_mla_decode_backend(self):
return self._create_trtllm_mla_decode_backend(backend="cute-dsl")
if not self.draft_model_runner.use_mla_backend:
raise ValueError(
"cutedsl_mla backend requires MLA model (use_mla_backend=True)."
)
from sglang.srt.layers.attention.cutedsl_mla_backend import (
CuteDslMLAMultiStepDraftBackend,
)
return CuteDslMLAMultiStepDraftBackend(
self.draft_model_runner, self.topk, self.speculative_num_steps
)
def _create_tokenspeed_mla_decode_backend(self):
if not self.draft_model_runner.use_mla_backend:
+25
View File
@@ -3608,6 +3608,31 @@ def require_mlp_sync(server_args: ServerArgs):
return server_args.enable_dp_attention or require_gathered_buffer(server_args)
def get_cuda_graph_batch_size_alignment(server_args: ServerArgs) -> int:
alignment = 1
if server_args.enable_two_batch_overlap:
alignment *= 2
if require_gathered_buffer(server_args):
alignment *= get_parallel().attn_tp_size
if alignment % get_parallel().attn_cp_size != 0:
alignment *= get_parallel().attn_cp_size
return alignment
def get_cuda_graph_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int:
return ceil_align(max_batch_size, get_cuda_graph_batch_size_alignment(server_args))
def get_eager_max_batch_size(server_args: ServerArgs, max_batch_size: int) -> int:
if not require_mlp_sync(server_args):
return max_batch_size
from sglang.srt.layers.cp.padding import get_cp_padding_align_size
max_batch_size = ceil_align(max_batch_size, get_parallel().attn_tp_size)
return ceil_align(max_batch_size, get_cp_padding_align_size())
def find_local_repo_dir(repo_id: str, revision: Optional[str] = None) -> Optional[str]:
import huggingface_hub as hf