[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
@@ -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):