[Kimi] Support DCP + DSpark (ported from kimi-k3 branch) (#32828)

This commit is contained in:
Khoa Pham
2026-07-31 17:39:00 -07:00
committed by GitHub
parent e4c4faf8a2
commit 1496bfee93
8 changed files with 445 additions and 17 deletions
@@ -80,14 +80,21 @@ _TOKENSPEED_MAX_Q_LEN = 8
def _get_tokenspeed_workspace(
device: torch.device, num_heads: int, kv_lora_rank: int
device: torch.device,
num_heads: int,
kv_lora_rank: int,
max_q_len: int = _TOKENSPEED_MAX_Q_LEN,
) -> torch.Tensor:
from sglang.srt.runtime_context import get_resources
# DCP target verification gathers Q to the full head count before launching
# TokenSpeed; size for that launch shape, not the rank-local head count.
num_heads *= get_parallel().attn_dcp_size
max_q_len = max(max_q_len, _TOKENSPEED_MAX_Q_LEN)
needed = (
tokenspeed_mla.get_num_sm(device)
* num_heads
* _TOKENSPEED_MAX_Q_LEN
* max_q_len
* (kv_lora_rank + 1)
* 4
)
@@ -133,7 +140,12 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
self._tokenspeed_workspace: Optional[torch.Tensor] = None
if is_tokenspeed_mla_available():
self._tokenspeed_workspace = _get_tokenspeed_workspace(
self.device, self.num_q_heads, self.kv_lora_rank
self.device,
self.num_q_heads,
self.kv_lora_rank,
max_q_len=(
model_runner.server_args.max_speculative_num_draft_tokens or 1
),
)
# Pre-JIT the prefill kernel variants. Each cute.compile takes 1-2
+6 -13
View File
@@ -16,12 +16,7 @@ from sglang.srt.hardware_backend.npu.dsv4.dsv4_common_hooks import (
from sglang.srt.mem_cache.allocator.swa import SWATokenToKVPoolAllocator
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
from sglang.srt.runtime_context import (
get_schedule,
get_server_args,
get_serving,
get_spec,
)
from sglang.srt.runtime_context import get_serving, get_spec
from sglang.srt.utils.common import ceil_align
if TYPE_CHECKING:
@@ -183,8 +178,8 @@ def release_kv_cache(req: Req, tree_cache: BasePrefixCache, is_insert: bool = Tr
def _release_overallocated_kv_indices(
req: Req, start_p: int, end_p: int, tree_cache: BasePrefixCache
) -> None:
global_server_args = get_server_args()
page_size = get_schedule().page_size
allocator = tree_cache.token_to_kv_pool_allocator
page_size = allocator.page_size
spec_algo = get_spec().speculative_algorithm
# strip_thinking_cache intentionally reports output tokens as overallocated
@@ -201,11 +196,9 @@ def _release_overallocated_kv_indices(
indices_to_free = tree_cache.req_to_token_pool.req_to_token[req.req_pool_idx][
start_p:end_p
]
# start_p is ceil-aligned above: never shares a page with
# cache_finished_req's tail frees in the same group.
tree_cache.token_to_kv_pool_allocator.free_segment(
indices_to_free, start_pos=start_p
)
# start_p is aligned to the allocator's physical page size above, so it
# never shares a page with cache_finished_req's tail free in this group.
allocator.free_segment(indices_to_free, start_pos=start_p)
def available_and_evictable_str(tree_cache: BasePrefixCache) -> str:
@@ -275,6 +275,16 @@ class KVCacheConfigurator:
full_max_total_num_tokens = config.full_max_total_num_tokens
swa_max_total_num_tokens = config.swa_max_total_num_tokens
# Draft pools are replicated, not DCP-sharded, yet consume the shared
# allocator's virtual locs in [0, max_total * dcp_size) untranslated.
dcp_size = self.server_args.dcp_size
if self.is_draft_worker and dcp_size > 1:
max_total_num_tokens *= dcp_size
if full_max_total_num_tokens is not None:
full_max_total_num_tokens *= dcp_size
if swa_max_total_num_tokens is not None:
swa_max_total_num_tokens *= dcp_size
# DSV4 compressed-attention pool sizes. Draft worker reuses target's
# full/swa sizes but does NOT own c4/c128/state pools (those live on
# the target rank only); zero them out regardless of what config holds.
@@ -177,7 +177,7 @@ class DefaultPoolConfigurator(MemoryPoolConfigurator):
self._cell_size = scale_kv_cell_size_per_token_for_dflash(
target_cell_size_per_token=self._cell_size,
target_num_layers=int(num_layers),
draft_num_layers=int(draft_num_layers),
draft_num_layers=int(draft_num_layers) * kvc.server_args.dcp_size,
)
def _compute_cell_size(self, kvc: KVCacheConfigurator, num_layers: int) -> int:
@@ -5,6 +5,7 @@ from typing import Optional
import torch
from sglang.srt.configs.hybrid_arch import mambaish_config
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import ScheduleBatch
@@ -57,6 +58,7 @@ from sglang.srt.speculative.spec_utils import (
GrammarTree,
build_grammar_vocab_mask,
draft_tp_context,
prepare_mamba_track_for_verify,
)
from sglang.srt.utils import get_available_gpu_memory, is_cuda
@@ -245,6 +247,7 @@ class DSparkWorkerV2(BaseSpecWorker):
)
self._forced_budget_frac: Optional[float] = None
self._need_mamba_verify_commit = False
self._observers = DsparkStepObservers(
planner=self._verify_planner,
@@ -296,6 +299,12 @@ class DSparkWorkerV2(BaseSpecWorker):
def init_attention_backends(self):
with self._draft_context():
self._draft_worker.init_attention_backends()
self._need_mamba_verify_commit = mambaish_config(
self.model_runner.model_config
) is not None and hasattr(
self.model_runner.attn_backend,
"update_mamba_state_after_mtp_verify",
)
def init_cuda_graphs(self):
capture_decode_cuda_graph = not get_exec().graph.disable_cuda_graph
@@ -587,6 +596,7 @@ class DSparkWorkerV2(BaseSpecWorker):
and self._simulate_acc_len <= 0
and not batch.has_grammar
)
prepare_mamba_track_for_verify(batch)
with self._observers.segment(InfoSegment.TARGET_VERIFY):
if run_compact:
target_verify, hidden_strided = self._verify_executor.run_compact(
@@ -644,6 +654,13 @@ class DSparkWorkerV2(BaseSpecWorker):
else:
on_publish(accept.new_seq_lens)
self._commit_target_mamba_states_after_verify(
batch=batch,
seq_lens_pre_verify=prefix_lens,
seq_lens_post_verify=accept.new_seq_lens,
commit_lens=accept.commit_lens,
)
folded_commit = folded_accept and epilogue.folds_commit
if not folded_commit:
self._verify_executor.commit_hidden(
@@ -699,5 +716,52 @@ class DSparkWorkerV2(BaseSpecWorker):
new_seq_lens=accept.new_seq_lens,
)
def _commit_target_mamba_states_after_verify(
self,
*,
batch: ScheduleBatch,
seq_lens_pre_verify: torch.Tensor,
seq_lens_post_verify: torch.Tensor,
commit_lens: torch.Tensor,
) -> None:
"""Commit the last accepted verify step's KDA/mamba state (chain
layout: step index = commit_lens - 1) into the persistent caches."""
if not self._need_mamba_verify_commit:
return
# Chain layout only: step index = commit_lens - 1. A tree (topk > 1)
# layout would need the accept-index mapping the shared spec_utils
# commit helper does.
assert self.server_args.speculative_eagle_topk in (None, 1)
attn_backend = self.target_worker.model_runner.attn_backend
last_correct_step_indices = commit_lens.to(torch.int64) - 1
mamba_steps_to_track = None
if batch.mamba_track_indices is not None:
mamba_track_interval = self.server_args.mamba_track_interval
to_track_mask = (
seq_lens_pre_verify // mamba_track_interval
!= seq_lens_post_verify // mamba_track_interval
)
tracking_point = (
seq_lens_post_verify // mamba_track_interval * mamba_track_interval
)
to_track_ith = torch.clamp(tracking_point - seq_lens_pre_verify - 1, min=0)
can_track_mask = to_track_mask & (
to_track_ith < commit_lens.to(to_track_ith.dtype)
)
mamba_steps_to_track = torch.where(
can_track_mask,
to_track_ith.to(torch.int64),
torch.full_like(to_track_ith, -1, dtype=torch.int64),
)
attn_backend.update_mamba_state_after_mtp_verify(
last_correct_step_indices=last_correct_step_indices,
mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track,
model=self.target_worker.model_runner.model,
)
def get_confidence_budget_prepare(self):
return self._verify_planner.confidence_budget_prepare()