[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()
@@ -0,0 +1,222 @@
import json
import socket
import tempfile
import time
import unittest
from pathlib import Path
from types import SimpleNamespace
from urllib.parse import urlparse
import requests
import torch
from sglang.srt.utils import kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.run_eval import run_eval
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
register_cuda_ci(est_time=350, stage="extra-b", runner_config="4-gpu-b200")
KIMI_LINEAR_MODEL = "moonshotai/Kimi-Linear-48B-A3B-Instruct"
GSM8K_SCORE_THRESHOLD = 0.88
CUDA_GRAPH_MAX_BS_DECODE = 128
MAX_RUNNING_REQUESTS = 128
GSM8K_NUM_THREADS = 128
def _has_four_blackwell_gpus() -> bool:
if not torch.cuda.is_available() or torch.cuda.device_count() < 4:
return False
return all(
torch.cuda.get_device_capability(device_index) >= (10, 0)
for device_index in range(4)
)
def _write_dummy_qwen3_dspark_draft(root: Path) -> str:
draft_dir = root / "qwen3-dspark-kimi-proxy"
draft_dir.mkdir()
config = {
"architectures": ["Qwen3DSparkModel"],
"model_type": "qwen3",
"dtype": "bfloat16",
"hidden_size": 2304,
"intermediate_size": 9216,
"num_hidden_layers": 5,
"num_attention_heads": 16,
"num_key_value_heads": 4,
"head_dim": 128,
"hidden_act": "silu",
"rms_norm_eps": 1e-5,
"attention_bias": False,
"attention_dropout": 0.0,
"max_position_embeddings": 1048576,
"rope_parameters": {
"rope_theta": 10000.0,
"rope_type": "default",
},
"vocab_size": 163840,
"bos_token_id": 163584,
"eos_token_id": 163586,
"mask_token_id": 163839,
"block_size": 7,
"markov_rank": 256,
"markov_head_type": "vanilla",
"enable_confidence_head": True,
"confidence_head_with_markov": True,
"num_target_layers": 27,
"target_layer_ids": [1, 7, 13, 19, 26],
"layer_types": ["full_attention"] * 5,
"tie_word_embeddings": False,
"use_cache": True,
}
(draft_dir / "config.json").write_text(json.dumps(config), encoding="utf-8")
return str(draft_dir)
def _wait_for_port_release(base_url: str, timeout: float = 30.0) -> None:
parsed = urlparse(base_url)
deadline = time.monotonic() + timeout
while time.monotonic() < deadline:
with socket.socket() as sock:
sock.settimeout(0.2)
if sock.connect_ex((parsed.hostname, parsed.port)) != 0:
return
time.sleep(0.1)
raise TimeoutError(f"Server port was not released after {timeout}s: {base_url}")
@unittest.skipUnless(
_has_four_blackwell_gpus(),
"Kimi Linear TokenSpeed DCP + DSpark requires four Blackwell GPUs",
)
class TestKimiLinearDCPDSpark4(CustomTestCase):
base_url = DEFAULT_URL_FOR_TEST
def _generate(self, prompts: list[str], *, max_new_tokens: int):
response = requests.post(
self.base_url + "/generate",
json={
"text": prompts,
"sampling_params": {
"temperature": 0,
"max_new_tokens": max_new_tokens,
"ignore_eos": True,
},
},
timeout=300,
)
response.raise_for_status()
outputs = response.json()
self.assertIsInstance(outputs, list)
self.assertEqual(len(outputs), len(prompts))
for output in outputs:
self.assertTrue(output["text"].strip())
self.assertGreater(output["meta_info"]["completion_tokens"], 0)
return outputs
def _run_static(self, *, draft_path: str, qrep: bool):
other_args = [
"--tp-size",
"4",
"--dcp-size",
"4",
"--max-running-requests",
str(MAX_RUNNING_REQUESTS),
"--attention-backend",
"tokenspeed_mla",
"--kv-cache-dtype",
"fp8_e4m3",
"--dcp-comm-backend",
"a2a",
"--speculative-algorithm",
"DSPARK",
"--speculative-draft-model-path",
draft_path,
"--speculative-draft-load-format",
"dummy",
"--speculative-attention-mode",
"decode",
"--speculative-draft-attention-backend",
"trtllm_mha",
"--cuda-graph-max-bs-decode",
str(CUDA_GRAPH_MAX_BS_DECODE),
"--cuda-graph-backend-prefill",
"disabled",
"--trust-remote-code",
"--random-seed",
"0",
"--dtype",
"bfloat16",
"--mem-fraction-static",
"0.80",
]
if qrep:
other_args.append("--dcp-replicate-q-proj")
process = popen_launch_server(
KIMI_LINEAR_MODEL,
self.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH * 8,
other_args=other_args,
env={
"SGLANG_PREP_IN_CUDA_GRAPH": "1",
"SGLANG_RAGGED_VERIFY_MODE": "static",
},
)
try:
captured_outputs = self._generate(
[
f"Reply with one short word for captured request {index}: the sky is"
for index in range(2)
],
max_new_tokens=8,
)
max_graph_outputs = self._generate(
[
f"Reply with one short word for graph request {index}: ice is"
for index in range(CUDA_GRAPH_MAX_BS_DECODE)
],
max_new_tokens=8,
)
requests.get(self.base_url + "/flush_cache", timeout=30).raise_for_status()
metrics = run_eval(
SimpleNamespace(
base_url=self.base_url,
model=KIMI_LINEAR_MODEL,
eval_name="gsm8k",
api="completion",
max_tokens=512,
num_examples=200,
num_threads=GSM8K_NUM_THREADS,
num_shots=5,
)
)
return captured_outputs + max_graph_outputs, float(metrics["score"])
finally:
kill_process_tree(process.pid, wait_timeout=60)
_wait_for_port_release(self.base_url)
def test_static_verify_cuda_graph(self):
with tempfile.TemporaryDirectory() as tmp:
root = Path(tmp)
draft_path = _write_dummy_qwen3_dspark_draft(root)
for qrep in (True, False):
with self.subTest(qrep=qrep):
outputs, score = self._run_static(draft_path=draft_path, qrep=qrep)
self.assertGreaterEqual(score, GSM8K_SCORE_THRESHOLD)
self.assertTrue(
all(
not output["meta_info"].get("spec_cap_lens_histogram")
for output in outputs
)
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,87 @@
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
from sglang.srt.layers.attention import tokenspeed_mla_backend as backend_module
from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend
from sglang.srt.layers.attention.trtllm_mla_backend import TRTLLMMLADecodeMetadata
from sglang.srt.layers.dcp.layout import get_dcp_lens
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=60, stage="base-b", runner_config="4-gpu-b200")
NUM_DRAFT_TOKENS = 8
DCP_SIZE = 4
DCP_RANK = 2
def _make_backend(bs: int):
backend = object.__new__(TokenspeedMLABackend)
backend.num_draft_tokens = NUM_DRAFT_TOKENS
metadata = TRTLLMMLADecodeMetadata(
block_kv_indices=torch.full((bs, 4), -1, dtype=torch.int32, device="cuda"),
seq_lens_k=torch.zeros(bs, dtype=torch.int32, device="cuda"),
global_seq_lens_k=torch.zeros(bs, dtype=torch.int32, device="cuda"),
)
backend.decode_cuda_graph_metadata = {bs: metadata}
return backend, metadata
def _apply(backend, *, bs: int, seq_lens: torch.Tensor, forward_mode):
parallel = SimpleNamespace(dcp_enabled=True, dcp_size=DCP_SIZE, dcp_rank=DCP_RANK)
with (
patch.object(backend_module, "get_parallel", return_value=parallel),
patch.object(backend, "_fill_dcp_block_kv_indices") as fill,
):
backend._apply_cuda_graph_metadata(
bs=bs,
req_pool_indices=torch.arange(bs, dtype=torch.int32, device="cuda"),
seq_lens=seq_lens,
forward_mode=forward_mode,
)
return fill
@unittest.skipUnless(torch.cuda.is_available(), "DCP metadata buffers live on CUDA")
class TestTokenspeedMLADCPMetadata(CustomTestCase):
def test_target_verify_splits_global_and_local_lengths(self):
bs = 3
backend, metadata = _make_backend(bs)
prefix_lens = torch.tensor([10, 20, 30], dtype=torch.int32, device="cuda")
fill = _apply(
backend,
bs=bs,
seq_lens=prefix_lens,
forward_mode=ForwardMode.TARGET_VERIFY,
)
expected_global = prefix_lens + NUM_DRAFT_TOKENS
expected_local = get_dcp_lens(expected_global, DCP_SIZE, DCP_RANK).to(
torch.int32
)
torch.testing.assert_close(metadata.global_seq_lens_k, expected_global)
torch.testing.assert_close(metadata.seq_lens_k, expected_local)
fill.assert_called_once()
torch.testing.assert_close(fill.call_args.args[2], expected_local)
def test_decode_does_not_add_draft_tokens(self):
bs = 3
backend, _ = _make_backend(bs)
seq_lens = torch.tensor([10, 20, 30], dtype=torch.int32, device="cuda")
fill = _apply(
backend, bs=bs, seq_lens=seq_lens, forward_mode=ForwardMode.DECODE
)
expected_local = get_dcp_lens(seq_lens, DCP_SIZE, DCP_RANK).to(torch.int32)
fill.assert_called_once()
torch.testing.assert_close(fill.call_args.args[2], expected_local)
if __name__ == "__main__":
unittest.main()
@@ -6,11 +6,14 @@ deferral. See PagedTokenToKVPoolAllocator.free_segment for why unique is avoided
"""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
from sglang.srt.mem_cache.allocator.paged import PagedTokenToKVPoolAllocator
from sglang.srt.mem_cache.common import _release_overallocated_kv_indices
from sglang.test.ci.ci_register import register_cpu_ci
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
@@ -103,6 +106,43 @@ class TestFreeSegment(unittest.TestCase):
with self.assertRaises(AssertionError):
alloc.free_group_end()
def test_overallocated_tail_uses_allocator_page_size_under_dcp(self):
# Scaled-down DCP example: the configured logical page is 1 while the
# allocator page is widened to 4. cache_finished_req has already freed
# the committed tail [4, 5), so over-allocation cleanup for [5, 7)
# must not release the same physical page again.
alloc = _make_allocator()
alloc.debug_mode = True
row = _make_kv_row(alloc, 2 * PAGE_SIZE)
tree_cache = SimpleNamespace(
token_to_kv_pool_allocator=alloc,
req_to_token_pool=SimpleNamespace(req_to_token=row.unsqueeze(0)),
)
req = SimpleNamespace(req_pool_idx=0)
before = len(alloc.free_pages)
alloc.free_group_begin()
alloc.free_segment(row[PAGE_SIZE : PAGE_SIZE + 1], start_pos=PAGE_SIZE)
with (
patch(
"sglang.srt.mem_cache.common.get_spec",
return_value=SimpleNamespace(speculative_algorithm="DSPARK"),
),
patch(
"sglang.srt.mem_cache.common.get_serving",
return_value=SimpleNamespace(strip_thinking_cache=False),
),
):
_release_overallocated_kv_indices(
req,
start_p=PAGE_SIZE + 1,
end_p=2 * PAGE_SIZE - 1,
tree_cache=tree_cache,
)
alloc.free_group_end()
self.assertEqual(len(alloc.free_pages), before + 1)
class TestFreeSegments(unittest.TestCase):
def _freed_by_segments(self, num_tokens, spans):