[Kimi] Support DCP + DSpark (ported from kimi-k3 branch) (#32828)
This commit is contained in:
@@ -80,14 +80,21 @@ _TOKENSPEED_MAX_Q_LEN = 8
|
|||||||
|
|
||||||
|
|
||||||
def _get_tokenspeed_workspace(
|
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:
|
) -> torch.Tensor:
|
||||||
from sglang.srt.runtime_context import get_resources
|
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 = (
|
needed = (
|
||||||
tokenspeed_mla.get_num_sm(device)
|
tokenspeed_mla.get_num_sm(device)
|
||||||
* num_heads
|
* num_heads
|
||||||
* _TOKENSPEED_MAX_Q_LEN
|
* max_q_len
|
||||||
* (kv_lora_rank + 1)
|
* (kv_lora_rank + 1)
|
||||||
* 4
|
* 4
|
||||||
)
|
)
|
||||||
@@ -133,7 +140,12 @@ class TokenspeedMLABackend(TRTLLMMLABackend):
|
|||||||
self._tokenspeed_workspace: Optional[torch.Tensor] = None
|
self._tokenspeed_workspace: Optional[torch.Tensor] = None
|
||||||
if is_tokenspeed_mla_available():
|
if is_tokenspeed_mla_available():
|
||||||
self._tokenspeed_workspace = _get_tokenspeed_workspace(
|
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
|
# Pre-JIT the prefill kernel variants. Each cute.compile takes 1-2
|
||||||
|
|||||||
@@ -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.allocator.swa import SWATokenToKVPoolAllocator
|
||||||
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
|
from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache, EvictParams
|
||||||
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
from sglang.srt.mem_cache.memory_pool import HybridReqToTokenPool, ReqToTokenPool
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import get_serving, get_spec
|
||||||
get_schedule,
|
|
||||||
get_server_args,
|
|
||||||
get_serving,
|
|
||||||
get_spec,
|
|
||||||
)
|
|
||||||
from sglang.srt.utils.common import ceil_align
|
from sglang.srt.utils.common import ceil_align
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
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(
|
def _release_overallocated_kv_indices(
|
||||||
req: Req, start_p: int, end_p: int, tree_cache: BasePrefixCache
|
req: Req, start_p: int, end_p: int, tree_cache: BasePrefixCache
|
||||||
) -> None:
|
) -> None:
|
||||||
global_server_args = get_server_args()
|
allocator = tree_cache.token_to_kv_pool_allocator
|
||||||
page_size = get_schedule().page_size
|
page_size = allocator.page_size
|
||||||
spec_algo = get_spec().speculative_algorithm
|
spec_algo = get_spec().speculative_algorithm
|
||||||
|
|
||||||
# strip_thinking_cache intentionally reports output tokens as overallocated
|
# 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][
|
indices_to_free = tree_cache.req_to_token_pool.req_to_token[req.req_pool_idx][
|
||||||
start_p:end_p
|
start_p:end_p
|
||||||
]
|
]
|
||||||
# start_p is ceil-aligned above: never shares a page with
|
# start_p is aligned to the allocator's physical page size above, so it
|
||||||
# cache_finished_req's tail frees in the same group.
|
# never shares a page with cache_finished_req's tail free in this group.
|
||||||
tree_cache.token_to_kv_pool_allocator.free_segment(
|
allocator.free_segment(indices_to_free, start_pos=start_p)
|
||||||
indices_to_free, start_pos=start_p
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def available_and_evictable_str(tree_cache: BasePrefixCache) -> str:
|
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
|
full_max_total_num_tokens = config.full_max_total_num_tokens
|
||||||
swa_max_total_num_tokens = config.swa_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
|
# DSV4 compressed-attention pool sizes. Draft worker reuses target's
|
||||||
# full/swa sizes but does NOT own c4/c128/state pools (those live on
|
# 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.
|
# 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(
|
self._cell_size = scale_kv_cell_size_per_token_for_dflash(
|
||||||
target_cell_size_per_token=self._cell_size,
|
target_cell_size_per_token=self._cell_size,
|
||||||
target_num_layers=int(num_layers),
|
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:
|
def _compute_cell_size(self, kvc: KVCacheConfigurator, num_layers: int) -> int:
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ from typing import Optional
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.configs.hybrid_arch import mambaish_config
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
@@ -57,6 +58,7 @@ from sglang.srt.speculative.spec_utils import (
|
|||||||
GrammarTree,
|
GrammarTree,
|
||||||
build_grammar_vocab_mask,
|
build_grammar_vocab_mask,
|
||||||
draft_tp_context,
|
draft_tp_context,
|
||||||
|
prepare_mamba_track_for_verify,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import get_available_gpu_memory, is_cuda
|
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._forced_budget_frac: Optional[float] = None
|
||||||
|
self._need_mamba_verify_commit = False
|
||||||
|
|
||||||
self._observers = DsparkStepObservers(
|
self._observers = DsparkStepObservers(
|
||||||
planner=self._verify_planner,
|
planner=self._verify_planner,
|
||||||
@@ -296,6 +299,12 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
def init_attention_backends(self):
|
def init_attention_backends(self):
|
||||||
with self._draft_context():
|
with self._draft_context():
|
||||||
self._draft_worker.init_attention_backends()
|
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):
|
def init_cuda_graphs(self):
|
||||||
capture_decode_cuda_graph = not get_exec().graph.disable_cuda_graph
|
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 self._simulate_acc_len <= 0
|
||||||
and not batch.has_grammar
|
and not batch.has_grammar
|
||||||
)
|
)
|
||||||
|
prepare_mamba_track_for_verify(batch)
|
||||||
with self._observers.segment(InfoSegment.TARGET_VERIFY):
|
with self._observers.segment(InfoSegment.TARGET_VERIFY):
|
||||||
if run_compact:
|
if run_compact:
|
||||||
target_verify, hidden_strided = self._verify_executor.run_compact(
|
target_verify, hidden_strided = self._verify_executor.run_compact(
|
||||||
@@ -644,6 +654,13 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
else:
|
else:
|
||||||
on_publish(accept.new_seq_lens)
|
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
|
folded_commit = folded_accept and epilogue.folds_commit
|
||||||
if not folded_commit:
|
if not folded_commit:
|
||||||
self._verify_executor.commit_hidden(
|
self._verify_executor.commit_hidden(
|
||||||
@@ -699,5 +716,52 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
new_seq_lens=accept.new_seq_lens,
|
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):
|
def get_confidence_budget_prepare(self):
|
||||||
return self._verify_planner.confidence_budget_prepare()
|
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
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
|
from sglang.srt.mem_cache.allocator.base import BaseTokenToKVPoolAllocator
|
||||||
from sglang.srt.mem_cache.allocator.paged import PagedTokenToKVPoolAllocator
|
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
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
|
||||||
@@ -103,6 +106,43 @@ class TestFreeSegment(unittest.TestCase):
|
|||||||
with self.assertRaises(AssertionError):
|
with self.assertRaises(AssertionError):
|
||||||
alloc.free_group_end()
|
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):
|
class TestFreeSegments(unittest.TestCase):
|
||||||
def _freed_by_segments(self, num_tokens, spans):
|
def _freed_by_segments(self, num_tokens, spans):
|
||||||
|
|||||||
Reference in New Issue
Block a user