[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(
|
||||
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
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user