[Spec] Add DSpark: confidence-scheduled speculative decoding (#30261)

Co-authored-by: sglang-bot <232288953+sglang-bot@users.noreply.github.com>
Co-authored-by: Claude Code <noreply@anthropic.com>
Co-authored-by: Codex <noreply@openai.com>
Co-authored-by: Liangsheng Yin <lsyincs@gmail.com>
Co-authored-by: Liangsheng Yin <hnyls2002@gmail.com>
This commit is contained in:
sglang-bot
2026-07-12 17:25:26 -05:00
committed by GitHub
co-authored by sglang-bot Claude Code Codex Liangsheng Yin Liangsheng Yin
parent 24d59d8d74
commit 6cc9352dfe
84 changed files with 17700 additions and 287 deletions
@@ -0,0 +1,94 @@
import unittest
from sglang.srt.utils import is_sm100_supported, kill_process_tree
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.kits.basic_api_contract_kit import BasicAPIContractMixin
from sglang.test.kits.basic_decode_correctness_kit import BasicDecodeCorrectnessMixin
from sglang.test.kits.basic_scheduler_stress_kit import BasicSchedulerStressMixin
from sglang.test.kits.eval_accuracy_kit import GSM8KMixin
from sglang.test.kits.fwd_occupancy_kit import FwdOccupancyMixin
from sglang.test.test_utils import (
DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
DEFAULT_URL_FOR_TEST,
CustomTestCase,
popen_launch_server,
)
register_cuda_ci(est_time=120, stage="base-b", runner_config="1-gpu-large")
TARGET_MODEL = "Qwen/Qwen3-14B"
DRAFT_MODEL = "deepseek-ai/dspark_qwen3_14b_block7"
# trtllm_mha prefill requires SM100 (Blackwell); use the Hopper-native pair elsewhere.
if is_sm100_supported():
ATTENTION_BACKEND = "trtllm_mha"
DRAFT_ATTENTION_BACKEND = "fa4"
else:
ATTENTION_BACKEND = "fa3"
DRAFT_ATTENTION_BACKEND = "fa3"
class TestBasicSanityDSpark(
BasicAPIContractMixin,
BasicDecodeCorrectnessMixin,
BasicSchedulerStressMixin,
FwdOccupancyMixin,
GSM8KMixin,
CustomTestCase,
):
served_model_name = TARGET_MODEL
model = TARGET_MODEL
fwd_occupancy_threshold = 60
fwd_occupancy_max_new_tokens = 4096
fwd_occupancy_acc_length_threshold: float = 2.0
gsm8k_num_questions = 200
gsm8k_accuracy_thres = 0.80
gsm8k_accept_length_thres = 2.0
attention_backend = ATTENTION_BACKEND
draft_attention_backend = DRAFT_ATTENTION_BACKEND
process = None
@classmethod
def setUpClass(cls):
cls.base_url = DEFAULT_URL_FOR_TEST
cls.process = popen_launch_server(
TARGET_MODEL,
cls.base_url,
timeout=DEFAULT_TIMEOUT_FOR_SERVER_LAUNCH,
other_args=[
"--trust-remote-code",
"--attention-backend",
cls.attention_backend,
"--speculative-draft-attention-backend",
cls.draft_attention_backend,
"--speculative-algorithm",
"DSPARK",
"--speculative-draft-model-path",
DRAFT_MODEL,
"--cuda-graph-max-bs-decode",
"4",
"--mem-fraction-static",
"0.7",
"--page-size",
"1",
"--enable-metrics",
"--disable-piecewise-cuda-graph",
],
env={
"SGLANG_ENABLE_METRICS_DEVICE_TIMER": "1",
"SGLANG_RAGGED_VERIFY_MODE": "compact",
},
)
@classmethod
def tearDownClass(cls):
if cls.process is not None:
kill_process_tree(cls.process.pid)
if __name__ == "__main__":
unittest.main()
@@ -52,6 +52,7 @@ def _make_sampling_info(batch_size, vocab_size, device="cuda"):
top_ks=torch.zeros(batch_size, device=device, dtype=torch.int32),
min_ps=torch.zeros(batch_size, device=device),
is_all_greedy=True,
is_any_greedy=True,
need_top_p_sampling=False,
need_top_k_sampling=False,
need_min_p_sampling=False,
@@ -14,17 +14,17 @@ class TestResolveMinFreeSlots(unittest.TestCase):
def test_unset_non_dflash_disables(self):
# Unset + not DFlash -> trigger stays disabled.
self.assertIsNone(resolve_min_free_slots(None, 512, is_dflash=False))
self.assertIsNone(resolve_min_free_slots(None, 512, is_dflash_family=False))
def test_unset_dflash_auto_enables(self):
# Unset + DFlash -> falls back to the legacy formula (full mapping).
self.assertEqual(resolve_min_free_slots(None, 512, is_dflash=True), 4)
self.assertEqual(resolve_min_free_slots(None, 8, is_dflash=True), 2)
self.assertEqual(resolve_min_free_slots(None, 512, is_dflash_family=True), 4)
self.assertEqual(resolve_min_free_slots(None, 8, is_dflash_family=True), 2)
def test_unset_dflash_small_cluster_disables(self):
# DFlash auto-default still respects the < 8 guard.
self.assertIsNone(resolve_min_free_slots(None, 7, is_dflash=True))
self.assertIsNone(resolve_min_free_slots(None, 0, is_dflash=True))
self.assertIsNone(resolve_min_free_slots(None, 7, is_dflash_family=True))
self.assertIsNone(resolve_min_free_slots(None, 0, is_dflash_family=True))
def test_le_one_disables(self):
# <= 1 can never batch, so it is a no-op.
@@ -47,7 +47,7 @@ class TestResolveMinFreeSlots(unittest.TestCase):
def test_user_value_overrides_dflash_default(self):
# An explicit user value wins over the DFlash auto-default.
self.assertEqual(resolve_min_free_slots(3, 512, is_dflash=True), 3)
self.assertEqual(resolve_min_free_slots(3, 512, is_dflash_family=True), 3)
class TestMinFreeSlotsDelayer(unittest.TestCase):
@@ -0,0 +1,579 @@
import json
import math
import tempfile
import unittest
from pathlib import Path
import torch
from sglang.srt.speculative.dspark_components.dspark_block_accept_estimator import (
BlockAcceptEstimateRecorder,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
_GAMMA = 3
_VOCAB = 8
class _FakeLayout:
def __init__(self, verify_lens: torch.Tensor):
self.verify_lens = verify_lens
def _reference_logprob(
logits_row: torch.Tensor, token: int, temperature: float
) -> float:
scaled = logits_row.to(torch.float32) / temperature
return float(scaled[token] - torch.logsumexp(scaled, dim=-1))
def _make_recorder(tmp_dir: str) -> tuple[BlockAcceptEstimateRecorder, Path]:
path = Path(tmp_dir) / "estimate.jsonl"
recorder = BlockAcceptEstimateRecorder(path=str(path), gamma=_GAMMA, device="cpu")
return recorder, path
class _FakeDelayed:
def __init__(self):
self._pending = None
def step(self, *, compute_on_device, postprocess_on_host):
if self._pending is not None:
result, post = self._pending
self._pending = None
if result is not None:
post(result)
result = compute_on_device()
self._pending = (result, postprocess_on_host) if result is not None else None
def _observe(
recorder: BlockAcceptEstimateRecorder,
*,
forward_ct: int,
rid: str,
drafts: list[int],
corrected_logits: torch.Tensor,
target_logits: torch.Tensor,
verify_len: int,
correct_len: int,
bonus: int,
seq_len: int,
cap_trim: int = 0,
temperature: float = 1.0,
) -> None:
recorder.observe_verify_step(
forward_ct=forward_ct,
rids=[rid],
draft_tokens=torch.tensor([drafts], dtype=torch.int64),
corrected_logits=corrected_logits.unsqueeze(0),
draft_temperatures=torch.tensor([temperature], dtype=torch.float32),
greedy_mask=torch.tensor([False]),
target_logits=target_logits,
target_temperatures=torch.tensor([[temperature]], dtype=torch.float32),
truncated_sampling_mask=None,
logits_adjustments_are_noop=True,
correct_len=torch.tensor([correct_len], dtype=torch.int32),
cap_trim_lens=torch.tensor([cap_trim], dtype=torch.int32),
bonus=torch.tensor([bonus], dtype=torch.int64),
prefix_lens=torch.tensor([seq_len], dtype=torch.int64),
layout=_FakeLayout(torch.tensor([verify_len], dtype=torch.int32)),
)
def _read_records(path: Path) -> list[dict]:
return [json.loads(line) for line in path.read_text().splitlines()]
class TestBlockAcceptEstimateRecorder(CustomTestCase):
def test_exact_block_when_rejected_inside_window(self):
with tempfile.TemporaryDirectory() as tmp:
recorder, path = _make_recorder(tmp)
corrected = torch.randn(_GAMMA, _VOCAB)
target = torch.randn((_GAMMA + 1), _VOCAB)
_observe(
recorder,
forward_ct=1,
rid="r0",
drafts=[1, 2, 3],
corrected_logits=corrected,
target_logits=target,
verify_len=3,
correct_len=1,
bonus=5,
seq_len=10,
)
recorder._file.flush()
records = _read_records(path)
self.assertEqual(len(records), 1)
self.assertEqual(records[0]["w"], 2)
self.assertEqual(records[0]["cl"], 1)
self.assertNotIn("q_lp", records[0])
self.assertNotIn("pg", records[0])
def test_censored_block_gathers_q_and_same_step_bonus_row_p(self):
with tempfile.TemporaryDirectory() as tmp:
recorder, path = _make_recorder(tmp)
corrected = torch.randn(_GAMMA, _VOCAB)
target = torch.randn((_GAMMA + 1), _VOCAB)
drafts = [1, 2, 3]
_observe(
recorder,
forward_ct=1,
rid="r0",
drafts=drafts,
corrected_logits=corrected,
target_logits=target,
verify_len=2,
correct_len=1,
bonus=2,
seq_len=10,
)
recorder._file.flush()
records = _read_records(path)
self.assertEqual(len(records), 1)
record = records[0]
self.assertEqual(record["w"], 1)
self.assertEqual(record["cl"], 1)
self.assertEqual(record["trimmed_tokens"], [2, 3])
self.assertEqual(len(record["q_lp"]), 2)
self.assertAlmostEqual(
record["q_lp"][0],
_reference_logprob(corrected[1], 2, 1.0),
places=4,
)
self.assertAlmostEqual(
record["q_lp"][1],
_reference_logprob(corrected[2], 3, 1.0),
places=4,
)
self.assertEqual(len(record["pg"]), 1)
src_fct, offset, p_lp, draft_token, realized_token = record["pg"][0]
self.assertEqual((src_fct, offset), (1, 2))
self.assertEqual(draft_token, 2)
self.assertEqual(realized_token, 2)
self.assertAlmostEqual(
p_lp, _reference_logprob(target[1], 2, 1.0), places=4
)
def test_pending_block_resolves_in_later_step(self):
with tempfile.TemporaryDirectory() as tmp:
recorder, path = _make_recorder(tmp)
corrected1 = torch.randn(_GAMMA, _VOCAB)
target1 = torch.randn((_GAMMA + 1), _VOCAB)
_observe(
recorder,
forward_ct=1,
rid="r0",
drafts=[1, 2, 3],
corrected_logits=corrected1,
target_logits=target1,
verify_len=2,
correct_len=1,
bonus=2,
seq_len=10,
)
corrected2 = torch.randn(_GAMMA, _VOCAB)
target2 = torch.randn((_GAMMA + 1), _VOCAB)
_observe(
recorder,
forward_ct=2,
rid="r0",
drafts=[3, 6, 7],
corrected_logits=corrected2,
target_logits=target2,
verify_len=4,
correct_len=3,
bonus=4,
seq_len=12,
)
recorder._file.flush()
records = _read_records(path)
self.assertEqual(len(records), 2)
step2 = records[1]
self.assertEqual(len(step2["pg"]), 1)
src_fct, offset, p_lp, draft_token, realized_token = step2["pg"][0]
self.assertEqual((src_fct, offset), (1, 3))
self.assertEqual(draft_token, 3)
self.assertEqual(realized_token, 3)
self.assertAlmostEqual(
p_lp, _reference_logprob(target2[0], 3, 1.0), places=4
)
self.assertEqual(recorder._states["r0"].pending, [])
def test_divergence_drops_block_after_final_gather(self):
with tempfile.TemporaryDirectory() as tmp:
recorder, path = _make_recorder(tmp)
corrected1 = torch.randn(_GAMMA, _VOCAB)
target1 = torch.randn((_GAMMA + 1), _VOCAB)
_observe(
recorder,
forward_ct=1,
rid="r0",
drafts=[1, 2, 3],
corrected_logits=corrected1,
target_logits=target1,
verify_len=2,
correct_len=1,
bonus=6,
seq_len=10,
)
recorder._file.flush()
records = _read_records(path)
src_fct, offset, p_lp, draft_token, realized_token = records[0]["pg"][0]
self.assertEqual(draft_token, 2)
self.assertEqual(realized_token, 6)
self.assertEqual(recorder._states["r0"].pending, [])
def test_temperature_scales_logprobs(self):
with tempfile.TemporaryDirectory() as tmp:
recorder, path = _make_recorder(tmp)
corrected = torch.randn(_GAMMA, _VOCAB)
target = torch.randn((_GAMMA + 1), _VOCAB)
_observe(
recorder,
forward_ct=1,
rid="r0",
drafts=[1, 2, 3],
corrected_logits=corrected,
target_logits=target,
verify_len=2,
correct_len=1,
bonus=2,
seq_len=10,
temperature=0.7,
)
recorder._file.flush()
record = _read_records(path)[0]
self.assertAlmostEqual(
record["q_lp"][0],
_reference_logprob(corrected[1], 2, 0.7),
places=4,
)
self.assertAlmostEqual(
record["pg"][0][2],
_reference_logprob(target[1], 2, 0.7),
places=4,
)
def test_greedy_row_is_skipped_but_seq_len_bookkeeping_advances(self):
with tempfile.TemporaryDirectory() as tmp:
recorder, path = _make_recorder(tmp)
corrected = torch.randn(_GAMMA, _VOCAB)
target = torch.randn((_GAMMA + 1), _VOCAB)
recorder.observe_verify_step(
forward_ct=1,
rids=["r0"],
draft_tokens=torch.tensor([[1, 2, 3]], dtype=torch.int64),
corrected_logits=corrected.unsqueeze(0),
draft_temperatures=torch.tensor([1.0]),
greedy_mask=torch.tensor([True]),
target_logits=target,
target_temperatures=torch.tensor([[1.0]]),
truncated_sampling_mask=None,
logits_adjustments_are_noop=True,
correct_len=torch.tensor([2], dtype=torch.int32),
cap_trim_lens=torch.tensor([0], dtype=torch.int32),
bonus=torch.tensor([5], dtype=torch.int64),
prefix_lens=torch.tensor([10], dtype=torch.int64),
layout=_FakeLayout(torch.tensor([4], dtype=torch.int32)),
)
recorder._file.flush()
self.assertEqual(path.read_text(), "")
self.assertEqual(recorder._states["r0"].expected_seq_len, 13)
def test_seq_len_discontinuity_drops_pending_blocks(self):
with tempfile.TemporaryDirectory() as tmp:
recorder, path = _make_recorder(tmp)
corrected1 = torch.randn(_GAMMA, _VOCAB)
target1 = torch.randn((_GAMMA + 1), _VOCAB)
_observe(
recorder,
forward_ct=1,
rid="r0",
drafts=[1, 2, 3],
corrected_logits=corrected1,
target_logits=target1,
verify_len=2,
correct_len=1,
bonus=2,
seq_len=10,
)
self.assertEqual(len(recorder._states["r0"].pending), 1)
corrected2 = torch.randn(_GAMMA, _VOCAB)
target2 = torch.randn((_GAMMA + 1), _VOCAB)
_observe(
recorder,
forward_ct=2,
rid="r0",
drafts=[3, 6, 7],
corrected_logits=corrected2,
target_logits=target2,
verify_len=4,
correct_len=3,
bonus=4,
seq_len=11,
)
recorder._file.flush()
records = _read_records(path)
self.assertNotIn("pg", records[1])
self.assertEqual(recorder._states["r0"].pending, [])
self.assertEqual(recorder._discontinuity_drop_ct, 1)
def test_truncated_sampling_row_is_excluded_while_clean_row_records(self):
with tempfile.TemporaryDirectory() as tmp:
recorder, path = _make_recorder(tmp)
corrected = torch.randn(2, _GAMMA, _VOCAB)
target = torch.randn(2 * (_GAMMA + 1), _VOCAB)
recorder.observe_verify_step(
forward_ct=1,
rids=["r_clean", "r_top_p"],
draft_tokens=torch.tensor([[1, 2, 3], [4, 5, 6]], dtype=torch.int64),
corrected_logits=corrected,
draft_temperatures=torch.tensor([1.0, 1.0]),
greedy_mask=torch.tensor([False, False]),
target_logits=target,
target_temperatures=torch.tensor([[1.0], [1.0]]),
truncated_sampling_mask=torch.tensor([False, True]),
logits_adjustments_are_noop=True,
correct_len=torch.tensor([1, 1], dtype=torch.int32),
cap_trim_lens=torch.tensor([0, 0], dtype=torch.int32),
bonus=torch.tensor([2, 7], dtype=torch.int64),
prefix_lens=torch.tensor([10, 20], dtype=torch.int64),
layout=_FakeLayout(torch.tensor([2, 2], dtype=torch.int32)),
)
recorder._file.flush()
records = _read_records(path)
self.assertEqual([r["rid"] for r in records], ["r_clean"])
self.assertEqual(recorder._skipped_step_ct, 0)
self.assertEqual(recorder._states["r_top_p"].pending, [])
self.assertEqual(recorder._states["r_top_p"].expected_seq_len, 22)
def _offline_estimate(path: Path, gamma: int) -> tuple[float, float, int]:
from collections import defaultdict
blocks: list[dict] = []
gathers: dict[tuple, list] = defaultdict(list)
for line in path.read_text().splitlines():
rec = json.loads(line)
blocks.append(rec)
for src_fct, offset, p_lp, draft_token, realized_token in rec.get("pg", []):
gathers[(rec["rid"], src_fct)].append(
[offset, p_lp, draft_token, realized_token]
)
los: list[float] = []
his: list[float] = []
for rec in blocks:
cl, w = rec["cl"], rec["w"]
if "q_lp" not in rec:
los.append(cl + 1.0)
his.append(cl + 1.0)
continue
q_lps = rec["q_lp"]
entries = {e[0]: e for e in gathers.get((rec["rid"], rec["fct"]), [])}
base, prod, lo_extra, tail = w + 1.0, 1.0, 0.0, 0.0
for offset in range(w + 1, gamma + 1):
entry = entries.get(offset)
if entry is None:
tail = prod * (gamma - offset + 1)
break
_, p_lp, draft_token, realized_token = entry
a = min(1.0, math.exp(p_lp - q_lps[offset - w - 1]))
prod *= a
lo_extra += prod
if draft_token != realized_token:
if offset < gamma:
tail = prod * (gamma - offset)
break
los.append(base + lo_extra)
his.append(base + lo_extra + tail)
n = len(los)
return sum(los) / n, sum(his) / n, n
class TestOnlineCeilingEstimate(CustomTestCase):
def test_online_estimate_matches_offline_aggregation(self):
bs, steps = 4, 14
gen = torch.Generator().manual_seed(11)
seq = [50 + 3 * b for b in range(bs)]
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "est.jsonl"
recorder = BlockAcceptEstimateRecorder(
path=str(path),
gamma=_GAMMA,
device="cpu",
online_log_interval=3,
online_window_steps=2,
)
for t in range(steps):
verify_lens, correct_lens, drafts, bonus, prefix = [], [], [], [], []
for b in range(bs):
vl = int(torch.randint(1, _GAMMA + 2, (1,), generator=gen))
window = vl - 1
cl = int(torch.randint(0, window + 1, (1,), generator=gen))
row = torch.randint(0, _VOCAB, (_GAMMA,), generator=gen).tolist()
if cl < _GAMMA and int(torch.randint(0, 2, (1,), generator=gen)):
bt = row[cl]
else:
bt = int(torch.randint(0, _VOCAB, (1,), generator=gen))
verify_lens.append(vl)
correct_lens.append(cl)
drafts.append(row)
bonus.append(bt)
prefix.append(seq[b])
seq[b] += cl + 1
recorder.observe_verify_step(
forward_ct=t + 1,
rids=[f"r{b}" for b in range(bs)],
draft_tokens=torch.tensor(drafts, dtype=torch.int64),
corrected_logits=torch.randn(bs, _GAMMA, _VOCAB, generator=gen),
draft_temperatures=torch.ones(bs),
greedy_mask=torch.zeros(bs, dtype=torch.bool),
target_logits=torch.randn(bs * (_GAMMA + 1), _VOCAB, generator=gen),
target_temperatures=torch.ones(bs),
truncated_sampling_mask=None,
logits_adjustments_are_noop=True,
correct_len=torch.tensor(correct_lens, dtype=torch.int32),
cap_trim_lens=torch.tensor(
[_GAMMA - (v - 1) for v in verify_lens], dtype=torch.int32
),
bonus=torch.tensor(bonus, dtype=torch.int64),
prefix_lens=torch.tensor(prefix, dtype=torch.int64),
layout=_FakeLayout(torch.tensor(verify_lens, dtype=torch.int32)),
)
recorder.drain_pending_online()
recorder._file.flush()
off_lo, off_hi, off_n = _offline_estimate(path, _GAMMA)
snap = recorder.online_estimate()
self.assertIsNotNone(snap)
self.assertEqual(snap.cumulative_blocks, off_n)
self.assertAlmostEqual(snap.cumulative_lo, off_lo, places=6)
self.assertAlmostEqual(snap.cumulative_hi, off_hi, places=6)
self.assertGreater(off_n, bs)
self.assertLessEqual(snap.window_blocks, snap.cumulative_blocks)
def test_online_window_evicts_forward_passes_outside_horizon(self):
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "est.jsonl"
recorder = BlockAcceptEstimateRecorder(
path=str(path),
gamma=_GAMMA,
device="cpu",
online_log_interval=1,
online_window_steps=3,
)
target = torch.randn((_GAMMA + 1), _VOCAB)
seq = 10
for t in range(10):
cl = t % 3
_observe(
recorder,
forward_ct=t + 1,
rid="r0",
drafts=[1, 2, 3],
corrected_logits=torch.randn(_GAMMA, _VOCAB),
target_logits=target,
verify_len=_GAMMA + 1,
correct_len=cl,
bonus=5,
seq_len=seq,
)
seq += cl + 1
snap = recorder.online_estimate()
self.assertIsNotNone(snap)
self.assertLessEqual(snap.window_horizon, 3)
self.assertLessEqual(snap.window_blocks, 3)
self.assertEqual(snap.cumulative_blocks, 10)
class TestNaturalStopEosTail(CustomTestCase):
def _finalize_kept_block(self, *, natural_stop: bool):
with tempfile.TemporaryDirectory() as tmp:
recorder, _ = _make_recorder(tmp)
_observe(
recorder,
forward_ct=1,
rid="r0",
drafts=[1, 2, 3],
corrected_logits=torch.randn(_GAMMA, _VOCAB),
target_logits=torch.randn((_GAMMA + 1), _VOCAB),
verify_len=2,
correct_len=1,
bonus=2,
seq_len=10,
)
self.assertEqual(len(recorder._states["r0"].pending), 1)
recorder.note_request_finished(rid="r0", natural_stop=natural_stop)
self.assertNotIn("r0", recorder._states)
return recorder.online_estimate()
def test_natural_eos_caps_tail_to_zero(self):
snap = self._finalize_kept_block(natural_stop=True)
self.assertEqual(snap.cumulative_blocks, 1)
self.assertAlmostEqual(snap.cumulative_lo, snap.cumulative_hi, places=6)
def test_external_finish_keeps_optimistic_tail(self):
snap = self._finalize_kept_block(natural_stop=False)
self.assertEqual(snap.cumulative_blocks, 1)
self.assertGreater(snap.cumulative_hi, snap.cumulative_lo)
class TestAsyncFinishIntent(CustomTestCase):
def test_intent_buffered_then_applied_at_next_drain(self):
with tempfile.TemporaryDirectory() as tmp:
recorder, _ = _make_recorder(tmp)
recorder._delayed = _FakeDelayed()
_observe(
recorder,
forward_ct=1,
rid="r0",
drafts=[1, 2, 3],
corrected_logits=torch.randn(_GAMMA, _VOCAB),
target_logits=torch.randn((_GAMMA + 1), _VOCAB),
verify_len=2,
correct_len=1,
bonus=2,
seq_len=10,
)
recorder.note_request_finished(rid="r0", natural_stop=True)
self.assertIn("r0", recorder._finish_intents)
self.assertIsNone(recorder.online_estimate())
_observe(
recorder,
forward_ct=2,
rid="r1",
drafts=[4, 5, 6],
corrected_logits=torch.randn(_GAMMA, _VOCAB),
target_logits=torch.randn((_GAMMA + 1), _VOCAB),
verify_len=_GAMMA + 1,
correct_len=1,
bonus=7,
seq_len=20,
)
self.assertNotIn("r0", recorder._finish_intents)
self.assertNotIn("r0", recorder._states)
snap = recorder.online_estimate()
self.assertEqual(snap.cumulative_blocks, 1)
self.assertAlmostEqual(snap.cumulative_lo, snap.cumulative_hi, places=6)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,159 @@
import unittest
import torch
from sglang.srt.environ import envs
from sglang.srt.speculative.dspark_components.dspark_observability import (
ConfidenceMetricsProbe,
PerPositionConfidenceMetrics,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
def _cpu_metrics(gamma: int) -> PerPositionConfidenceMetrics:
return PerPositionConfidenceMetrics(gamma=gamma, device=torch.device("cpu"))
class TestPerPositionConfidenceMetrics(CustomTestCase):
def test_perfectly_calibrated_has_low_ece(self):
torch.manual_seed(0)
n = 40000
survival = torch.full((n, 1), 0.3, dtype=torch.float64)
prefix_mask = (torch.rand(n, 1) < 0.3).to(torch.float64)
metrics = _cpu_metrics(gamma=1)
metrics.update(survival=survival, prefix_mask=prefix_mask)
row = metrics.compute()[0]
self.assertLess(row["ece"], 0.03)
self.assertAlmostEqual(row["pred_mean"], 0.3, places=4)
def test_overconfident_has_high_ece_and_pred_above_target(self):
torch.manual_seed(0)
n = 40000
survival = torch.full((n, 1), 0.9, dtype=torch.float64)
prefix_mask = (torch.rand(n, 1) < 0.3).to(torch.float64)
metrics = _cpu_metrics(gamma=1)
metrics.update(survival=survival, prefix_mask=prefix_mask)
row = metrics.compute()[0]
self.assertGreater(row["ece"], 0.4)
self.assertGreater(row["pred_mean"], row["target_mean"])
def test_separable_scores_give_auc_near_one(self):
torch.manual_seed(0)
n = 20000
pos = torch.rand(n, 1) * 0.3 + 0.7
neg = torch.rand(n, 1) * 0.3
survival = torch.cat([pos, neg], dim=0)
prefix_mask = torch.cat([torch.ones(n, 1), torch.zeros(n, 1)], dim=0)
metrics = _cpu_metrics(gamma=1)
metrics.update(survival=survival, prefix_mask=prefix_mask)
self.assertGreater(metrics.compute()[0]["auc"], 0.99)
def test_random_scores_give_auc_near_half(self):
torch.manual_seed(0)
n = 40000
survival = torch.rand(n, 1)
prefix_mask = (torch.rand(n, 1) < 0.5).to(torch.float64)
metrics = _cpu_metrics(gamma=1)
metrics.update(survival=survival, prefix_mask=prefix_mask)
auc = metrics.compute()[0]["auc"]
self.assertGreater(auc, 0.45)
self.assertLess(auc, 0.55)
def test_batched_update_matches_per_sample_update(self):
torch.manual_seed(0)
bs, gamma = 32, 5
survival = torch.rand(bs, gamma)
prefix_mask = (torch.rand(bs, gamma) < 0.5).to(torch.float64)
batched = _cpu_metrics(gamma=gamma)
batched.update(survival=survival, prefix_mask=prefix_mask)
per_sample = _cpu_metrics(gamma=gamma)
for row_idx in range(bs):
per_sample.update(
survival=survival[row_idx : row_idx + 1],
prefix_mask=prefix_mask[row_idx : row_idx + 1],
)
for name in (
"coarse_count",
"coarse_pred",
"coarse_target",
"fine_pos",
"fine_neg",
"brier_num",
):
self.assertTrue(
torch.allclose(
getattr(batched, name), getattr(per_sample, name), atol=1e-9
),
msg=name,
)
def _probe_inputs(bs: int, gamma: int):
torch.manual_seed(0)
verify_num_draft_tokens = gamma + 1
vocab = 16
verify_ids_2d = torch.randint(0, vocab, (bs, verify_num_draft_tokens))
target_logits = torch.randn(bs * verify_num_draft_tokens, vocab)
confidence_raw = torch.randn(bs, gamma)
return verify_ids_2d, target_logits, confidence_raw
class TestConfidenceMetricsProbe(CustomTestCase):
def _observe(self, probe, *, carries_confidence=True, is_compact_mode=False):
verify_ids_2d, target_logits, confidence_raw = _probe_inputs(bs=3, gamma=4)
probe.maybe_observe(
carries_confidence=carries_confidence,
is_compact_mode=is_compact_mode,
confidence_raw=confidence_raw,
verify_ids_2d=verify_ids_2d,
target_logits=target_logits,
bs=3,
)
def test_disabled_env_is_noop(self):
probe = ConfidenceMetricsProbe(gamma=4, verify_num_draft_tokens=5, tp_rank=0)
self._observe(probe)
self.assertIsNone(probe._metrics)
self.assertEqual(probe._step_ct, 0)
def test_non_rank0_is_noop(self):
probe = ConfidenceMetricsProbe(gamma=4, verify_num_draft_tokens=5, tp_rank=1)
with envs.SGLANG_DSPARK_DEBUG_CONFIDENCE_METRICS.override(True):
self._observe(probe)
self.assertIsNone(probe._metrics)
def test_missing_confidence_head_is_noop(self):
probe = ConfidenceMetricsProbe(gamma=4, verify_num_draft_tokens=5, tp_rank=0)
with envs.SGLANG_DSPARK_DEBUG_CONFIDENCE_METRICS.override(True):
self._observe(probe, carries_confidence=False)
self.assertIsNone(probe._metrics)
def test_compact_mode_warns_once_and_skips(self):
probe = ConfidenceMetricsProbe(gamma=4, verify_num_draft_tokens=5, tp_rank=0)
with envs.SGLANG_DSPARK_DEBUG_CONFIDENCE_METRICS.override(True):
self._observe(probe, is_compact_mode=True)
self.assertTrue(probe._compact_warned)
self._observe(probe, is_compact_mode=True)
self.assertIsNone(probe._metrics)
self.assertEqual(probe._step_ct, 0)
def test_enabled_path_accumulates(self):
probe = ConfidenceMetricsProbe(
gamma=4, verify_num_draft_tokens=5, tp_rank=0, print_every=2
)
with envs.SGLANG_DSPARK_DEBUG_CONFIDENCE_METRICS.override(True):
self._observe(probe)
self.assertIsInstance(probe._metrics, PerPositionConfidenceMetrics)
self.assertEqual(probe._step_ct, 1)
self._observe(probe)
self.assertEqual(probe._step_ct, 2)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,99 @@
import random
import unittest
from sglang.srt.speculative.dspark_components.dspark_planner import (
dp_global_verify_tier_num_tokens,
local_verify_tier_num_tokens,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
class TestLocalVerifyTierNumTokens(CustomTestCase):
def test_no_budget_returns_sentinel(self):
self.assertEqual(
local_verify_tier_num_tokens(
bs=8,
verify_token_budget=None,
verify_num_draft_tokens=6,
min_verify_len=1,
),
-1,
)
def test_budget_adds_to_anchor_floor(self):
self.assertEqual(
local_verify_tier_num_tokens(
bs=8,
verify_token_budget=10,
verify_num_draft_tokens=6,
min_verify_len=1,
),
18,
)
# Clamp/floor variants (verify-all clamp, min_verify_len floor, min=0) are
# covered by the TestBusyIdleGraphKeyIdentity sweep bounds.
class TestDpGlobalVerifyTierNumTokens(CustomTestCase):
def test_any_sentinel_pins_everyone(self):
# The sweep never emits a -1 contribution, so this is the only guard
# on "any rank without a budget pins everyone"; losing it forks graph
# keys across DP ranks.
self.assertIsNone(
dp_global_verify_tier_num_tokens(global_tier_num_tokens=[100, -1, 50, 0])
)
class TestBusyIdleGraphKeyIdentity(CustomTestCase):
def test_busy_and_idle_floors_agree_on_random_topologies(self):
rng = random.Random(20260703)
for _ in range(2000):
verify_num_draft_tokens = rng.randint(2, 8)
min_verify_len = rng.randint(0, verify_num_draft_tokens - 1)
effective_min = max(min_verify_len, 1)
num_ranks = rng.randint(1, 8)
contributions = []
num_reqs_per_rank = []
for _ in range(num_ranks):
if rng.random() < 0.3:
num_reqs_per_rank.append(0)
contributions.append(0)
continue
bs = rng.randint(1, 512)
budget = rng.randint(0, bs * verify_num_draft_tokens)
num_reqs_per_rank.append(bs)
contributions.append(
local_verify_tier_num_tokens(
bs=bs,
verify_token_budget=budget,
verify_num_draft_tokens=verify_num_draft_tokens,
min_verify_len=min_verify_len,
)
)
tier_num_tokens = dp_global_verify_tier_num_tokens(
global_tier_num_tokens=contributions
)
global_num_reqs = max(num_reqs_per_rank)
if tier_num_tokens is None:
self.assertEqual(global_num_reqs, 0)
continue
self.assertGreaterEqual(tier_num_tokens, global_num_reqs * effective_min)
self.assertLessEqual(
tier_num_tokens, global_num_reqs * verify_num_draft_tokens
)
busy_floor = min(tier_num_tokens, global_num_reqs * verify_num_draft_tokens)
self.assertEqual(busy_floor, tier_num_tokens)
idle_lens_total = global_num_reqs
idle_bucket_input = max(idle_lens_total, tier_num_tokens)
self.assertEqual(idle_bucket_input, tier_num_tokens)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,88 @@
import unittest
from types import SimpleNamespace
from sglang.srt.arg_groups.speculative_hook import (
_handle_dspark,
_target_checkpoint_bundles_dspark_draft,
)
from sglang.srt.server_args import ServerArgs
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
_BUNDLED_MODEL_PATH = "deepseek-ai/DeepSeek-V4-Flash-DSpark"
_PLAIN_MODEL_PATH = "deepseek-ai/DeepSeek-V4-Flash"
def _bundled_hf_config() -> SimpleNamespace:
return SimpleNamespace(
architectures=["DeepseekV4ForCausalLM"],
dspark_block_size=5,
dspark_markov_rank=256,
dspark_target_layer_ids=[40, 41, 42],
dspark_noise_token_id=128799,
)
def _plain_hf_config() -> SimpleNamespace:
return SimpleNamespace(architectures=["DeepseekV4ForCausalLM"])
def _make_dspark_server_args(
*, model_path: str, hf_config: SimpleNamespace
) -> ServerArgs:
server_args = ServerArgs(model_path="dummy")
server_args.model_path = model_path
server_args.device = "cuda"
server_args.speculative_algorithm = "DSPARK"
server_args.speculative_draft_model_path = None
server_args.speculative_dspark_block_size = 5
server_args.model_config = SimpleNamespace(hf_config=hf_config)
return server_args
class TestTargetCheckpointBundlesDsparkDraft(CustomTestCase):
def test_bundled_dsv4_config_is_detected(self):
server_args = _make_dspark_server_args(
model_path=_BUNDLED_MODEL_PATH, hf_config=_bundled_hf_config()
)
self.assertTrue(_target_checkpoint_bundles_dspark_draft(server_args))
def test_plain_target_config_is_not_detected(self):
server_args = _make_dspark_server_args(
model_path=_PLAIN_MODEL_PATH, hf_config=_plain_hf_config()
)
self.assertFalse(_target_checkpoint_bundles_dspark_draft(server_args))
class TestDsparkDraftPathDefaulting(CustomTestCase):
def test_bundled_checkpoint_defaults_draft_path_to_model_path(self):
server_args = _make_dspark_server_args(
model_path=_BUNDLED_MODEL_PATH, hf_config=_bundled_hf_config()
)
_handle_dspark(server_args)
self.assertEqual(server_args.speculative_draft_model_path, _BUNDLED_MODEL_PATH)
self.assertEqual(server_args.speculative_num_draft_tokens, 6)
def test_plain_target_without_draft_path_raises(self):
server_args = _make_dspark_server_args(
model_path=_PLAIN_MODEL_PATH, hf_config=_plain_hf_config()
)
with self.assertRaises(ValueError):
_handle_dspark(server_args)
def test_explicit_draft_path_is_not_overwritten(self):
server_args = _make_dspark_server_args(
model_path=_BUNDLED_MODEL_PATH, hf_config=_bundled_hf_config()
)
server_args.speculative_draft_model_path = "deepseek-ai/some-other-dspark-draft"
_handle_dspark(server_args)
self.assertEqual(
server_args.speculative_draft_model_path,
"deepseek-ai/some-other-dspark-draft",
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,398 @@
import unittest
import torch
from sglang.srt.environ import envs
from sglang.srt.speculative.dspark_components.dspark_observability import (
DecodeStepObservation,
DsparkInfoDumper,
InfoComponent,
_PendingStep,
logger,
resolve_components,
resolve_enabled_components,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
class FakeClock:
def __init__(self) -> None:
self.now = 100.0
def __call__(self) -> float:
return self.now
def advance(self, seconds: float) -> None:
self.now += seconds
def make_dumper(components, **kwargs):
clock = FakeClock()
dumper = DsparkInfoDumper(
components=set(components),
gamma=5,
verify_num_draft_tokens=6,
attn_tp_rank=0,
device=torch.device("cpu"),
mode_value="static",
clock=clock,
**kwargs,
)
return dumper, clock
def make_obs(
*,
forward_ct,
bs=4,
num_verify_tokens=24,
predicted_step_ms=None,
predicted_theta=None,
):
return DecodeStepObservation(
forward_ct=forward_ct,
bs=bs,
mode="static",
budget=100,
lag_steps=0,
num_verify_tokens=num_verify_tokens,
verify_tokens_local=num_verify_tokens,
verify_tokens_dp_synced=num_verify_tokens,
verify_tokens_graph_key=num_verify_tokens,
predicted_step_ms=predicted_step_ms,
predicted_theta=predicted_theta,
verify_lens=torch.full((bs,), 6, dtype=torch.int32),
confidence=torch.full((bs, 5), 0.9),
req_pool_indices=torch.arange(bs, dtype=torch.int64),
prefix_lens=torch.full((bs,), 128, dtype=torch.int64),
draft_tokens=torch.zeros((bs, 5), dtype=torch.int64),
bonus_tokens=torch.zeros((bs,), dtype=torch.int64),
correct_len=torch.full((bs,), 3, dtype=torch.int32),
cap_trim_lens=torch.zeros((bs,), dtype=torch.int32),
commit_lens=torch.full((bs,), 4, dtype=torch.int32),
rids=[f"r{i}" for i in range(bs)],
)
class TestResolveComponents(CustomTestCase):
def test_empty_disables(self):
self.assertEqual(resolve_components(()), set())
def test_all_expands_to_every_component(self):
self.assertEqual(resolve_components(("all",)), set(InfoComponent))
def test_subset_and_whitespace_are_kept(self):
self.assertEqual(
resolve_components((" core ", "reqs")),
{InfoComponent.CORE, InfoComponent.REQS},
)
def test_unknown_component_raises(self):
with self.assertRaises(ValueError):
resolve_components(("core", "bogus"))
def test_sps_record_env_enables_core_and_cpu_timing(self):
"""SGLANG_DSPARK_ENABLE_SPS_RECORD=1 is the published SPS-profiling
switch; it must keep enabling the components the table fit reads."""
with envs.SGLANG_DSPARK_ENABLE_SPS_RECORD.override(True):
self.assertEqual(
resolve_enabled_components(),
{InfoComponent.CORE, InfoComponent.STEP_CPU_TIME},
)
def test_sps_record_env_unions_with_debug_dump(self):
with envs.SGLANG_DSPARK_ENABLE_SPS_RECORD.override(True):
with envs.SGLANG_DSPARK_DEBUG_DUMP.override("reqs"):
self.assertEqual(
resolve_enabled_components(),
{
InfoComponent.CORE,
InfoComponent.STEP_CPU_TIME,
InfoComponent.REQS,
},
)
class TestCoreAndCpuTiming(CustomTestCase):
def test_disabled_dumper_records_nothing(self):
dumper, clock = make_dumper(set())
dumper.begin_step()
dumper.observe_decode_step(make_obs(forward_ct=1))
self.assertIsNone(dumper.dump())
def test_non_root_rank_is_disabled(self):
clock = FakeClock()
dumper = DsparkInfoDumper(
components={"core"},
gamma=5,
verify_num_draft_tokens=6,
attn_tp_rank=1,
device=torch.device("cpu"),
mode_value="static",
clock=clock,
)
self.assertFalse(dumper.enabled)
dumper.observe_decode_step(make_obs(forward_ct=1))
self.assertIsNone(dumper.dump())
def test_one_record_per_step_including_the_last(self):
dumper, clock = make_dumper({"core", "step_cpu_time"})
for forward_ct in range(1, 4):
dumper.observe_decode_step(make_obs(forward_ct=forward_ct))
clock.advance(0.01)
records = dumper.dump()["records"]
self.assertEqual([r["forward_ct"] for r in records], [1, 2, 3])
def test_step_cpu_ms_is_attributed_to_the_step_it_measures(self):
dumper, clock = make_dumper({"core", "step_cpu_time"})
dumper.observe_decode_step(make_obs(forward_ct=1))
clock.advance(0.02)
dumper.observe_decode_step(make_obs(forward_ct=2))
records = dumper.dump()["records"]
first = next(r for r in records if r["forward_ct"] == 1)
second = next(r for r in records if r["forward_ct"] == 2)
self.assertNotIn("step_cpu_ms", first)
self.assertAlmostEqual(second["step_cpu_ms"], 20.0, places=3)
def test_core_fields_present(self):
dumper, _ = make_dumper({"core"})
dumper.observe_decode_step(make_obs(forward_ct=7, bs=3, num_verify_tokens=18))
record = dumper.dump()["records"][0]
self.assertEqual(record["bs"], 3)
self.assertEqual(record["num_running_reqs"], 3)
self.assertEqual(record["num_verify_tokens"], 18)
self.assertEqual(record["mode"], "static")
def test_core_only_omits_timing_fields(self):
dumper, clock = make_dumper({"core"})
dumper.observe_decode_step(make_obs(forward_ct=1))
clock.advance(0.01)
dumper.observe_decode_step(make_obs(forward_ct=2))
for record in dumper.dump()["records"]:
self.assertNotIn("step_cpu_ms", record)
def test_non_decode_step_resets_cpu_pairing(self):
dumper, clock = make_dumper({"core", "step_cpu_time"})
dumper.observe_decode_step(make_obs(forward_ct=1))
clock.advance(0.02)
dumper.note_non_decode_step()
clock.advance(0.02)
dumper.observe_decode_step(make_obs(forward_ct=3))
records = dumper.dump()["records"]
self.assertEqual([r["forward_ct"] for r in records], [1, 3])
for record in records:
self.assertNotIn("step_cpu_ms", record)
def test_oversized_gap_nulls_cpu_ms_but_keeps_record(self):
dumper, clock = make_dumper({"core", "step_cpu_time"}, max_step_cpu_seconds=0.5)
dumper.observe_decode_step(make_obs(forward_ct=1))
clock.advance(0.6)
dumper.observe_decode_step(make_obs(forward_ct=2))
records = dumper.dump()["records"]
self.assertEqual([r["forward_ct"] for r in records], [1, 2])
second = next(r for r in records if r["forward_ct"] == 2)
self.assertNotIn("step_cpu_ms", second)
def test_ring_buffer_evicts_oldest(self):
dumper, clock = make_dumper({"core"}, max_records=3)
for forward_ct in range(1, 8):
dumper.observe_decode_step(make_obs(forward_ct=forward_ct))
clock.advance(0.01)
records = dumper.dump()["records"]
self.assertEqual([r["forward_ct"] for r in records], [5, 6, 7])
def test_dump_is_repeatable(self):
dumper, clock = make_dumper({"core"})
dumper.observe_decode_step(make_obs(forward_ct=1))
clock.advance(0.01)
dumper.observe_decode_step(make_obs(forward_ct=2))
self.assertEqual(dumper.dump(), dumper.dump())
def test_clear_drops_all_records_and_pending(self):
dumper, clock = make_dumper({"core"})
dumper.observe_decode_step(make_obs(forward_ct=1))
clock.advance(0.01)
dumper.observe_decode_step(make_obs(forward_ct=2))
dumper.clear()
self.assertEqual(dumper.dump()["records"], [])
dumper.observe_decode_step(make_obs(forward_ct=9))
clock.advance(0.01)
dumper.observe_decode_step(make_obs(forward_ct=10))
self.assertEqual([r["forward_ct"] for r in dumper.dump()["records"]], [9, 10])
class TestPredictedStepFields(CustomTestCase):
def test_predicted_fields_recorded_under_core(self):
dumper, clock = make_dumper({"core"})
dumper.observe_decode_step(
make_obs(forward_ct=1, predicted_step_ms=1.5, predicted_theta=200.0)
)
clock.advance(0.01)
dumper.observe_decode_step(make_obs(forward_ct=2))
record = next(r for r in dumper.dump()["records"] if r["forward_ct"] == 1)
self.assertAlmostEqual(record["predicted_step_ms"], 1.5)
self.assertAlmostEqual(record["predicted_theta"], 200.0)
def test_predicted_fields_omitted_when_none(self):
dumper, clock = make_dumper({"core"})
dumper.observe_decode_step(make_obs(forward_ct=1))
clock.advance(0.01)
dumper.observe_decode_step(make_obs(forward_ct=2))
record = next(r for r in dumper.dump()["records"] if r["forward_ct"] == 1)
self.assertNotIn("predicted_step_ms", record)
self.assertNotIn("predicted_theta", record)
def _pending(*, bs, budget, num_verify_tokens, predicted_step_ms):
return _PendingStep(
forward_ct=1,
bs=bs,
mode="compact",
budget=budget,
lag_steps=1,
num_verify_tokens=num_verify_tokens,
verify_tokens_local=num_verify_tokens,
verify_tokens_dp_synced=num_verify_tokens,
verify_tokens_graph_key=num_verify_tokens,
predicted_step_ms=predicted_step_ms,
predicted_theta=1.0,
step_cpu_ms=None,
rids=None,
future=None,
segment_events={},
)
class TestOnlineSpsReporter(CustomTestCase):
def test_report_interval_enables_dumper_and_gpu_timing(self):
dumper, _ = make_dumper(set(), sps_report_interval=2)
self.assertTrue(dumper.enabled)
self.assertIn(InfoComponent.STEP_GPU_TIME, dumper._components)
def test_report_interval_zero_leaves_dumper_disabled(self):
dumper, _ = make_dumper(set(), sps_report_interval=0)
self.assertFalse(dumper.enabled)
def test_reporter_logs_summary_every_interval_matched_steps(self):
dumper, _ = make_dumper(set(), sps_report_interval=2)
matched = dict(bs=4, budget=20, num_verify_tokens=24)
with self.assertLogs(logger, level="INFO") as cm:
dumper._report_sps_prediction(
pending=_pending(**matched, predicted_step_ms=10.0), step_gpu_ms=12.0
)
dumper._report_sps_prediction(
pending=_pending(**matched, predicted_step_ms=8.0), step_gpu_ms=9.0
)
self.assertEqual(sum("SPS prediction" in m for m in cm.output), 1)
self.assertEqual(dumper._sps_window, [])
def test_reporter_counts_mismatch_and_excludes_it_from_means(self):
dumper, _ = make_dumper(set(), sps_report_interval=1)
with self.assertLogs(logger, level="INFO") as cm:
dumper._report_sps_prediction(
pending=_pending(
bs=4, budget=99, num_verify_tokens=24, predicted_step_ms=10.0
),
step_gpu_ms=12.0,
)
dumper._report_sps_prediction(
pending=_pending(
bs=4, budget=20, num_verify_tokens=24, predicted_step_ms=10.0
),
step_gpu_ms=12.0,
)
self.assertTrue(any("M_mismatch_rate=50.0%" in m for m in cm.output))
def test_reporter_skips_steps_missing_prediction_or_actual(self):
dumper, _ = make_dumper(set(), sps_report_interval=2)
dumper._report_sps_prediction(
pending=_pending(
bs=4, budget=20, num_verify_tokens=24, predicted_step_ms=None
),
step_gpu_ms=12.0,
)
dumper._report_sps_prediction(
pending=_pending(
bs=4, budget=20, num_verify_tokens=24, predicted_step_ms=10.0
),
step_gpu_ms=None,
)
self.assertEqual(dumper._sps_window, [])
self.assertEqual(dumper._sps_mismatched, 0)
@unittest.skipUnless(torch.cuda.is_available(), "requires CUDA for d2h staging")
class TestReqsAndGpuTiming(CustomTestCase):
def _cuda_obs(self, *, forward_ct, bs=4):
obs = make_obs(forward_ct=forward_ct, bs=bs)
return DecodeStepObservation(
forward_ct=obs.forward_ct,
bs=obs.bs,
mode=obs.mode,
budget=obs.budget,
lag_steps=obs.lag_steps,
num_verify_tokens=obs.num_verify_tokens,
verify_tokens_local=obs.verify_tokens_local,
verify_tokens_dp_synced=obs.verify_tokens_dp_synced,
verify_tokens_graph_key=obs.verify_tokens_graph_key,
predicted_step_ms=obs.predicted_step_ms,
predicted_theta=obs.predicted_theta,
verify_lens=obs.verify_lens.cuda(),
confidence=obs.confidence.cuda(),
req_pool_indices=obs.req_pool_indices.cuda(),
prefix_lens=obs.prefix_lens.cuda(),
draft_tokens=obs.draft_tokens.cuda(),
bonus_tokens=obs.bonus_tokens.cuda(),
correct_len=obs.correct_len.cuda(),
cap_trim_lens=obs.cap_trim_lens.cuda(),
commit_lens=obs.commit_lens.cuda(),
rids=obs.rids,
)
def _make(self, components):
return DsparkInfoDumper(
components=set(components),
gamma=5,
verify_num_draft_tokens=6,
attn_tp_rank=0,
device=torch.device("cuda"),
mode_value="static",
)
def test_reqs_component_stages_per_request_detail(self):
dumper = self._make({"core", "reqs"})
dumper.observe_decode_step(self._cuda_obs(forward_ct=1, bs=3))
dumper.observe_decode_step(self._cuda_obs(forward_ct=2, bs=3))
record = next(r for r in dumper.dump()["records"] if r["forward_ct"] == 1)
self.assertEqual(len(record["reqs"]), 3)
req = record["reqs"][0]
self.assertEqual(req["rid"], "r0")
self.assertEqual(req["verify_len"], 6)
self.assertEqual(req["acc_len"], 4)
self.assertEqual(req["correct_drafts"], 3)
self.assertEqual(len(req["survival"]), 5)
def test_gpu_timing_populates_segment_fields(self):
dumper = self._make(
{"step_gpu_time", "draft_gpu_time", "target_verify_gpu_time"}
)
for forward_ct in (1, 2):
dumper.begin_step()
with dumper.segment("draft"):
torch.zeros(1024, device="cuda").sum()
with dumper.segment("target_verify"):
torch.zeros(1024, device="cuda").sum()
dumper.observe_decode_step(self._cuda_obs(forward_ct=forward_ct))
record = next(r for r in dumper.dump()["records"] if r["forward_ct"] == 1)
# The segments launch real kernels, so resolved event pairs must
# measure strictly positive time; 0.0 would mean the events never ran.
self.assertGreater(record["step_gpu_ms"], 0.0)
self.assertGreater(record["draft_gpu_ms"], 0.0)
self.assertGreater(record["target_verify_gpu_ms"], 0.0)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,562 @@
"""Seeded triton-vs-torch parity sweep for the DSpark kernels.
Guards against toolchain drift (triton/torch upgrades) silently diverging
the triton implementations from their torch references. Every method calls
the production kernel pair directly via `Cls.torch(...)` / `Cls.triton(...)`
(no env-var dispatch) on a small set of adversarial inputs and compares
exactly, or with the tolerance the kernel is specified to meet.
"""
import types
import unittest
import torch
from sglang.srt.layers.attention.dsv4 import attn_metadata_kernels
from sglang.srt.speculative import ragged_verify_kernels
from sglang.srt.speculative.dspark_components.dspark_planner import (
DSparkScheduleConfig,
)
from sglang.srt.speculative.dspark_components.kernels import (
dspark_accept,
dspark_attn_metadata,
dspark_draft_model,
dspark_schedule,
dspark_verify_window,
)
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=30, stage="base-b", runner_config="1-gpu-small")
DEVICE = torch.device("cuda")
VOCAB = 129280
def _ri(lo, hi, shape, dtype=torch.int64, g=None):
return torch.randint(lo, hi, shape, device=DEVICE, dtype=dtype, generator=g)
def _layout(verify_lens, graph_num_tokens):
return RaggedVerifyLayout.from_verify_lens_device(
verify_lens=verify_lens, graph_num_tokens=graph_num_tokens
)
class _Bf16Linear(torch.nn.Module):
quant_method = None
def __init__(self, weight):
super().__init__()
self.weight = weight
def forward(self, x):
return torch.nn.functional.linear(x, self.weight), None
def _case_accept_greedy(tc):
torch.manual_seed(0)
bs, t = 8, 6
candidates = _ri(0, 200, (bs, t))
target_logits = torch.randn(bs * t, 200, device=DEVICE)
for cutoff in (None, _ri(1, t + 1, (bs,), torch.int32)):
tc._parity(
dspark_accept.AcceptGreedy,
candidates=candidates,
target_logits=target_logits,
verify_num_draft_tokens=t,
cutoff_verify_lens=cutoff,
)
# gather_row_bonus: bonus token at a per-row column index.
table, idx = _ri(0, VOCAB, (64, t)), _ri(0, t, (64,), torch.int32)
ref = table[torch.arange(64, device=DEVICE), idx.long()]
tc._eq(dspark_accept.gather_row_bonus_triton(table=table, idx=idx), ref)
def _case_accept_sampling(tc):
torch.manual_seed(1)
bs, t = 64, 6
accept_index = _ri(0, bs * t, (bs, t))
predicts = _ri(0, VOCAB, (bs * t,))
correct_len = _ri(0, t, (bs,), torch.int32)
rows = torch.arange(bs, device=DEVICE)
ref = predicts[accept_index[rows, correct_len.long()].long()]
got = dspark_accept.gather_two_level_bonus_triton(
accept_index=accept_index, predicts=predicts, correct_len=correct_len
)
tc._eq(got, ref)
def _case_build_block_seq_lens_causal(tc):
torch.manual_seed(2)
seq_lens = _ri(1, 100000, (128,))
for block_size in (1, 5, 7):
tc._parity(
dspark_attn_metadata.BuildBlockSeqLensCausal,
seq_lens=seq_lens,
block_size=block_size,
device=DEVICE,
)
def _case_build_out_tokens(tc):
torch.manual_seed(3)
bs, gamma = 64, 5
for cl_dtype in (torch.int32, torch.int64):
# Bonus insertion swept through every position 0..gamma.
cl = (torch.arange(bs, device=DEVICE) % (gamma + 1)).to(cl_dtype)
tc._parity(
dspark_verify_window.BuildOutTokens,
draft_tokens=_ri(0, VOCAB, (bs, gamma)),
correct_len=cl,
bonus=_ri(0, VOCAB, (bs,)),
verify_num_draft_tokens=gamma + 1,
gamma=gamma,
)
def _case_build_ragged_verify_window(tc):
torch.manual_seed(4)
gamma, t, bs = 5, 6, 8
verify_lens = _ri(1, t + 1, (bs,), torch.int32)
batch = types.SimpleNamespace(
seq_lens=_ri(1, 20, (bs,)),
req_pool_indices=torch.randperm(bs + 3, device=DEVICE)[:bs],
)
model_runner = types.SimpleNamespace(
req_to_token_pool=types.SimpleNamespace(
req_to_token=_ri(0, 1_000_000, (bs + 3, 64), torch.int32)
)
)
for graph_num_tokens in (bs * t, (bs + 3) * t): # tight and bucket padding
tc._parity(
dspark_verify_window.BuildRaggedVerifyWindow,
batch=batch,
layout=_layout(verify_lens, graph_num_tokens),
draft_block_ids=_ri(0, VOCAB, (bs, gamma)),
draft_tokens=_ri(0, VOCAB, (bs, gamma)),
bs=bs,
device=DEVICE,
verify_num_draft_tokens=t,
model_runner=model_runner,
)
def _case_build_step_local(tc):
torch.manual_seed(5)
for org_width, per_partition, bias_dtype in (
(32320, 32384, torch.bfloat16),
(5000, 8192, torch.float32),
):
bias = (torch.randn(3, org_width, device=DEVICE) * 3.0).to(bias_dtype)
base = torch.randn(3, per_partition, device=DEVICE)
got, _ = tc._parity(
dspark_draft_model.BuildStepLocal, bias=bias, base_local=base
)
# Padding columns beyond org_width must stay pure base.
tc.assertTrue(torch.equal(got[:, org_width:], base[:, org_width:]))
def _case_cap_correct_len(tc):
torch.manual_seed(6)
bs, nd = 64, 6
verify_lens = _ri(1, nd + 1, (bs,), torch.int32)
for cl_dtype in (torch.int32, torch.int64):
cl = (torch.arange(bs, device=DEVICE) % (nd + 1)).to(cl_dtype)
tc._parity(dspark_accept.CapCorrectLen, correct_len=cl, verify_lens=verify_lens)
def _case_causal_swa_page_indices(tc):
swa, num_pool, pool_len, num_q = 128, 64, 600, 40
g = torch.Generator(device=DEVICE).manual_seed(7)
kw = dict(
req_to_token=_ri(0, 40000, (num_pool, pool_len), torch.int32, g),
full_to_swa_mapping=_ri(0, 1 << 20, (40000,), torch.int64, g),
req_pool_indices_repeated=_ri(0, num_pool, (num_q,), torch.int32, g),
swa_window=swa,
page_index_aligned_size=96,
)
# Lens short of / straddling / beyond the SWA window boundary.
for lo, hi in ((1, swa), (swa - 4, swa + 4), (swa + 1, pool_len)):
lens = _ri(lo, hi, (num_q,), torch.int32, g)
cls = attn_metadata_kernels.BuildCausalSwaPageIndices
ref = cls.torch(seq_lens_casual=lens, **kw)
got = cls.triton(seq_lens_casual=lens, **kw)
tc.assertEqual(got.shape, ref.shape)
tc.assertEqual(got.dtype, ref.dtype)
# Parity holds on the attended region; padding slots must be -1.
col = torch.arange(ref.shape[1], device=DEVICE).view(1, -1)
attended = col < torch.clamp(lens, max=swa).view(-1, 1)
tc.assertTrue(torch.equal(got[attended], ref[attended]))
tc.assertTrue(bool((got[~attended] == -1).all()))
def _case_commit_inject_layout(tc):
stride, num_pool, pool_len, n_full, bs = 7, 300, 400, 50000, 64
g = torch.Generator(device=DEVICE).manual_seed(8)
pool_perm = torch.randperm(num_pool, device=DEVICE, generator=g)
kw = dict(
req_pool_indices=pool_perm[:bs],
req_to_token=_ri(0, n_full, (num_pool, pool_len), torch.int64, g),
prefix_lens=_ri(1, pool_len - stride, (bs,), torch.int64, g),
block_pos_offsets=torch.arange(stride, device=DEVICE),
full_to_swa_mapping=_ri(0, 1 << 20, (n_full,), torch.int64, g),
commit_lens=_ri(0, stride + 1, (bs,), torch.int32, g),
stride=stride,
)
tc._parity(dspark_verify_window.BuildCommitInjectLayout, **kw)
# commit_len edges: 0 masks the whole row to -1, stride keeps it all.
kw.update(
req_pool_indices=kw["req_pool_indices"][:2],
prefix_lens=kw["prefix_lens"][:2],
commit_lens=torch.tensor([0, stride], device=DEVICE, dtype=torch.int32),
)
edge = dspark_verify_window.BuildCommitInjectLayout.triton(**kw)
swa_2d = edge.swa_loc.view(2, stride)
tc.assertTrue(bool((swa_2d[0] == -1).all()))
tc.assertTrue(bool((swa_2d[1] >= 0).all()))
def _case_commit_kv_proj(tc):
hidden, head_dim, num_stages = 1024, 576, 3
g = torch.Generator(device=DEVICE).manual_seed(9)
linears = [
_Bf16Linear(
(torch.randn(head_dim, hidden, device=DEVICE, generator=g) * 0.02).to(
torch.bfloat16
)
)
for _ in range(num_stages)
]
main_x = (torch.randn(56, hidden, device=DEVICE, generator=g) * 0.5).to(
torch.bfloat16
)
cls = dspark_draft_model.CommitKvProj
ref = cls.torch(main_x=main_x, wkv_linears=linears)
got = cls.triton(main_x=main_x, wkv_linears=linears)
tc.assertEqual(len(got), num_stages)
for kv_got, kv_ref in zip(got, ref):
tc.assertEqual(kv_got.shape, kv_ref.shape)
tc.assertTrue(kv_got.is_contiguous())
torch.testing.assert_close(kv_got.float(), kv_ref.float(), rtol=2e-2, atol=2e-3)
# fp8 blockwise weight dequant path (2x3 grid of 128x128 blocks).
out_dim, in_dim, block = 192, 384, 128
w8 = torch.randn(out_dim, in_dim, device=DEVICE, generator=g).to(
torch.float8_e4m3fn
)
scale = torch.rand(2, 3, device=DEVICE, generator=g) + 0.5
sf = scale.repeat_interleave(block, 0)[:out_dim]
sf = sf.repeat_interleave(block, 1)[:, :in_dim]
expected = (w8.to(torch.float32) * sf).to(torch.bfloat16)
stub = types.SimpleNamespace(weight=w8, weight_scale_inv=scale)
tc._eq(dspark_draft_model._dequant_linear_weight(stub), expected)
def _case_compact_layout(tc):
torch.manual_seed(10)
gamma, t, bs = 5, 6, 64
verify_lens = _ri(1, t + 1, (bs,), torch.int32)
total = int(verify_lens.sum().item())
for padded_total in (total, bs * t): # exact and bucket padding
tc._parity(
dspark_verify_window.CompactRowIndex,
verify_lens=verify_lens,
padded_total=padded_total,
device=DEVICE,
)
tc._parity(
dspark_verify_window.CompactVerifyIds,
draft_block_ids=_ri(0, VOCAB, (bs, gamma)),
draft_tokens=_ri(0, VOCAB, (bs, gamma)),
layout=_layout(verify_lens, padded_total),
device=DEVICE,
)
def _case_swa_page_indices(tc):
torch.manual_seed(11)
block_size, num_q, max_reqs, n_full = 5, 320, 300, 50000
_, gather = tc._parity(
dspark_attn_metadata.ComputeDsparkWindowGather,
seq_lens_casual=_ri(1, 300, (num_q,), torch.int32),
req_pool_indices_repeated=_ri(0, max_reqs, (num_q,)),
block_size=block_size,
swa_window=128,
)
tc._parity(
dspark_attn_metadata.BuildDsparkSwaPageIndices,
req_to_token=_ri(0, n_full, (max_reqs, 400), torch.int32),
full_to_swa_mapping=_ri(0, 20000, (n_full,), torch.int32),
req_pool_indices_per_request=gather.req_pool_indices_per_request,
offsets=gather.offsets,
invalid=gather.invalid,
out_loc=_ri(0, n_full, (num_q,)),
context_lens=gather.context_lens,
block_size=block_size,
swa_window=128,
page_index_aligned_size=64,
)
def _case_expand_prefill_causally(tc):
torch.manual_seed(12)
# Vectorized branch: ragged extends with padded token count.
bs = 64
extend = _ri(1, 8, (bs,))
num_tokens = int(extend.sum())
req_pool_indices = torch.randperm(512, device=DEVICE)[:bs]
seq_lens = _ri(8, 500, (bs,))
tc._parity(
attn_metadata_kernels.ExpandPrefillCausally,
req_pool_indices=req_pool_indices,
seq_lens=seq_lens,
extend_seq_lens=extend,
extend_start_loc=torch.cumsum(extend, dim=0) - extend,
seq_lens_cpu=None,
extend_seq_lens_cpu=None,
num_tokens=num_tokens,
padded_num_tokens=num_tokens + 5,
)
# Loop branch: uniform extend with CPU lens and no padding.
bs2, block = 8, 6
tc._parity(
attn_metadata_kernels.ExpandPrefillCausally,
req_pool_indices=req_pool_indices[:bs2],
seq_lens=seq_lens[:bs2],
extend_seq_lens=torch.full((bs2,), block, device=DEVICE),
extend_start_loc=None,
seq_lens_cpu=[int(x) for x in seq_lens[:bs2].tolist()],
extend_seq_lens_cpu=[block] * bs2,
num_tokens=bs2 * block,
padded_num_tokens=None,
)
def _case_finalize_accept_lens(tc):
torch.manual_seed(13)
bs = 64
for prefix_dtype in (torch.int32, torch.int64):
tc._parity(
dspark_accept.FinalizeAcceptLens,
correct_len=_ri(0, 7, (bs,), torch.int32),
cap_trim_lens=_ri(0, 4, (bs,)),
prefix_lens=_ri(1, 4000, (bs,), prefix_dtype),
)
def _case_mixed_accept_select(tc):
torch.manual_seed(14)
bs = 64
# Mixed dtypes between the greedy and sampling lanes.
tc._parity(
dspark_accept.SelectMixedAccept,
greedy_mask=torch.rand(bs, device=DEVICE) < 0.5,
greedy_len=_ri(0, 7, (bs,)),
greedy_bonus=_ri(0, 100000, (bs,)),
greedy_trim=_ri(0, 4, (bs,)),
sampling_len=_ri(0, 7, (bs,), torch.int32),
sampling_bonus=_ri(0, 100000, (bs,)),
sampling_trim=_ri(0, 4, (bs,), torch.int32),
)
def _case_padded_to_bucket(tc):
torch.manual_seed(15)
for bs, padded_bs, graph_num_tokens in ((3, 6, 16), (2, 8, 16), (8, 128, 768)):
verify_lens = _ri(1, 7, (bs,), torch.int32)
if int(verify_lens.sum()) > graph_num_tokens:
verify_lens = torch.ones(bs, dtype=torch.int32, device=DEVICE)
got, _ = tc._parity(
ragged_verify_kernels.PaddedToBucket,
verify_lens=verify_lens,
graph_num_tokens=graph_num_tokens,
bs=bs,
padded_bs=padded_bs,
)
# Padding rows must absorb exactly the leftover budget.
tc.assertEqual(int(got.to(torch.int64).sum()), graph_num_tokens)
if padded_bs > bs:
tc.assertTrue(torch.equal(got[:bs], verify_lens))
def _case_page_table_positions(tc):
num_pool, pool_len = 128, 4096
g = torch.Generator(device=DEVICE).manual_seed(16)
req_to_token = _ri(0, 1 << 20, (num_pool, pool_len), torch.int32, g)
# Large page + non-pool-aligned max_seq_len, then page_size 1.
for num_q, page_size, max_seq_len in ((300, 64, 4000), (56, 1, 4096)):
tc._parity(
attn_metadata_kernels.BuildPageTablePositions,
req_to_token=req_to_token,
req_pool_indices_repeated=_ri(0, num_pool, (num_q,), torch.int32, g),
seq_lens_casual=_ri(1, pool_len, (num_q,), torch.int64, g),
max_seq_len=max_seq_len,
page_size=page_size,
swa_window=128,
)
def _case_qo_indptr(tc):
torch.manual_seed(17)
cls = ragged_verify_kernels.BuildQoIndptr
for dtype in (torch.int32, torch.int64):
verify_lens = _ri(1, 8, (129,), dtype) # straddles the 128 block
ref = cls.torch(verify_lens=verify_lens)
got = cls.triton(verify_lens=verify_lens.to(torch.int32))
tc._eq(got, ref)
# Aliasing regression: the two outputs must not share storage.
vl = torch.tensor([3, 1, 5], device=DEVICE, dtype=torch.int32)
got = cls.triton(verify_lens=vl)
got.extend_start_loc.fill_(-7)
tc.assertEqual(got.qo_indptr[:2].tolist(), [0, 3])
def _case_sample_step_tokens(tc):
torch.manual_seed(18)
cls = dspark_draft_model.SampleStepTokens
# Injected noise makes stochastic sampling exactly comparable.
for vocab, dtype in ((130000, torch.bfloat16), (5003, torch.float32)):
bs = 3
tc._parity(
cls,
step_logits=(torch.randn(bs, vocab, device=DEVICE) * 4.0).to(dtype),
temperatures=torch.rand(bs, device=DEVICE) + 0.5,
greedy_mask=(torch.arange(bs, device=DEVICE) % 2) == 0,
exp_noise=torch.empty(bs, vocab, device=DEVICE).exponential_(1),
)
# Greedy tie straddling a triton block boundary picks the smaller index.
logits = torch.zeros(1, 2050, device=DEVICE)
logits[0, 1000] = logits[0, 1100] = 5.0
tokens = cls.triton(
step_logits=logits,
temperatures=torch.tensor([1.0], device=DEVICE),
greedy_mask=torch.tensor([True], device=DEVICE),
exp_noise=torch.ones(1, 2050, device=DEVICE),
)
tc.assertEqual(tokens.item(), 1000)
# Non-contiguous strided cropped view must match its contiguous copy.
view = (torch.randn(2, 129536, device=DEVICE) * 4.0)[:, :VOCAB]
tc.assertFalse(view.is_contiguous())
kw = dict(
temperatures=torch.rand(2, device=DEVICE) + 0.5,
greedy_mask=torch.tensor([True, False], device=DEVICE),
exp_noise=torch.empty(2, VOCAB, device=DEVICE).exponential_(1),
)
tc._eq(
cls.triton(step_logits=view, **kw),
cls.triton(step_logits=view.contiguous(), **kw),
)
def _case_scatter_compact_to_strided(tc):
torch.manual_seed(19)
t, bs, dim = 6, 8, 4096
verify_lens = _ri(1, t + 1, (bs,), torch.int32)
total = int(verify_lens.sum().item())
for graph_num_tokens in (total, bs * t): # exact and bucket padding
compact = torch.randn(
graph_num_tokens, dim, dtype=torch.bfloat16, device=DEVICE
)
tc._parity(
dspark_verify_window.ScatterCompactToStrided,
compact=compact,
layout=_layout(verify_lens, graph_num_tokens),
fill_value=0.0,
verify_num_draft_tokens=t,
)
def _case_schedule_verify_lens_topk(tc):
torch.manual_seed(20)
gamma, bs = 5, 64
cfg = DSparkScheduleConfig(gamma=gamma)
cls = dspark_schedule.ScheduleVerifyLensTopk
base = torch.rand(bs, gamma, device=DEVICE)
confidences = (
torch.full((bs, gamma), 0.5, device=DEVICE), # all-ties
(base * 4).floor() / 4, # coarse quantization
torch.where(base < 0.3, torch.zeros_like(base), base), # invalid zeros
)
for confidence in confidences:
for budget in (0, 1, 3, 7, 1000):
tc._parity(cls, confidence=confidence, budget=budget, cfg=cfg)
def _case_softmax_temp(tc):
g = torch.Generator(device=DEVICE).manual_seed(21)
cls = dspark_accept.SoftmaxTemp
# bf16 logits, non-power-of-two rows_per_request, full vocab.
logits = (torch.randn(56, VOCAB, device=DEVICE, generator=g) * 8.0).to(
torch.bfloat16
)
temps = (torch.rand(8, device=DEVICE, generator=g) * 1.5 + 0.05).float()
ref = cls.torch(logits=logits, temperatures=temps, rows_per_request=7)
got = cls.triton(logits=logits, temperatures=temps, rows_per_request=7)
tc.assertEqual(got.dtype, torch.float32)
torch.testing.assert_close(got, ref, rtol=1e-4, atol=1e-6)
torch.testing.assert_close(
got.sum(dim=-1), torch.ones_like(got.sum(dim=-1)), rtol=1e-5, atol=1e-5
)
# Column-shaped (bs, 1) temperatures.
logits2 = torch.randn(6, 512, device=DEVICE, generator=g).to(torch.bfloat16)
temps2 = (torch.rand(2, 1, device=DEVICE, generator=g) + 0.3).float()
ref2 = cls.torch(logits=logits2, temperatures=temps2, rows_per_request=3)
got2 = cls.triton(logits=logits2, temperatures=temps2, rows_per_request=3)
torch.testing.assert_close(got2, ref2, rtol=1e-5, atol=1e-7)
_CASES = [
("accept_greedy", _case_accept_greedy),
("accept_sampling", _case_accept_sampling),
("build_block_seq_lens_causal", _case_build_block_seq_lens_causal),
("build_out_tokens", _case_build_out_tokens),
("build_ragged_verify_window", _case_build_ragged_verify_window),
("build_step_local", _case_build_step_local),
("cap_correct_len", _case_cap_correct_len),
("causal_swa_page_indices", _case_causal_swa_page_indices),
("commit_inject_layout", _case_commit_inject_layout),
("commit_kv_proj", _case_commit_kv_proj),
("compact_layout", _case_compact_layout),
("swa_page_indices", _case_swa_page_indices),
("expand_prefill_causally", _case_expand_prefill_causally),
("finalize_accept_lens", _case_finalize_accept_lens),
("mixed_accept_select", _case_mixed_accept_select),
("padded_to_bucket", _case_padded_to_bucket),
("page_table_positions", _case_page_table_positions),
("qo_indptr", _case_qo_indptr),
("sample_step_tokens", _case_sample_step_tokens),
("scatter_compact_to_strided", _case_scatter_compact_to_strided),
("schedule_verify_lens_topk", _case_schedule_verify_lens_topk),
("softmax_temp", _case_softmax_temp),
]
class TestDsparkKernelParity(CustomTestCase):
def _eq(self, got, ref):
"""Exact comparison of tensors, tuples, and msgspec result structs."""
if isinstance(ref, tuple):
for g, r in zip(got, ref):
self._eq(g, r)
elif hasattr(ref, "__struct_fields__"):
for name in ref.__struct_fields__:
self._eq(getattr(got, name), getattr(ref, name))
elif isinstance(ref, torch.Tensor):
self.assertEqual(got.dtype, ref.dtype)
self.assertTrue(torch.equal(got, ref))
else:
self.assertEqual(got, ref)
def _parity(self, cls, **kw):
got, ref = cls.triton(**kw), cls.torch(**kw)
self._eq(got, ref)
return got, ref
def test_all_kernels_triton_matches_torch(self):
for name, case in _CASES:
with self.subTest(kernel=name):
case(self)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,558 @@
import functools
import types
import unittest
import torch
from sglang.srt.speculative.dspark_components.dspark_planner import (
DSparkScheduleConfig,
HostConfidenceBudgetPlanner,
VerifyBudgetDecision,
compute_verify_token_budget,
graph_tier_fill_budget,
)
from sglang.srt.speculative.dspark_components.dspark_sps import (
SpsAdditiveCostTable,
SpsCostTable,
)
from sglang.srt.speculative.dspark_components.kernels.dspark_schedule import (
schedule_verify_lens_topk_from_survival,
)
from sglang.srt.speculative.ragged_verify import RaggedVerifyLayout
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
def _flat_table(
steps_per_sec: float = 1.0, max_batch_tokens: int = 4096
) -> SpsCostTable:
return SpsCostTable(
sample_batch_tokens=[1],
sample_steps_per_sec=[steps_per_sec],
max_batch_tokens=max_batch_tokens,
)
def _cliff_table() -> SpsCostTable:
return SpsCostTable(
sample_batch_tokens=[1, 2, 3, 4, 5, 6, 7, 8],
sample_steps_per_sec=[1.0, 1.0, 1.0, 0.5, 0.45, 0.44, 0.43, 0.42],
max_batch_tokens=64,
)
def _additive_table() -> SpsAdditiveCostTable:
return SpsAdditiveCostTable(
bias_seconds=0.01,
bs_probes=[1, 100],
alpha_seconds=[0.0, 0.05],
m_probes=[1, 200],
theta_seconds=[0.0, 0.02],
)
def _survival_from_confidence(confidence: torch.Tensor) -> torch.Tensor:
return torch.cumprod(confidence, dim=1)
def _bruteforce_budget(
*,
history_survival_probs: torch.Tensor,
sps_table: SpsCostTable,
cfg: DSparkScheduleConfig,
) -> int:
num_requests = history_survival_probs.shape[0]
max_len = cfg.resolved_max_verify_len()
candidates = history_survival_probs[:, :max_len].flatten()
candidates = [float(x) for x in candidates.tolist() if float(x) >= cfg.survival_eps]
candidates.sort(reverse=True)
best_extra, best_theta = 0, float("-inf")
for extra in range(len(candidates) + 1):
tau_star = num_requests + sum(candidates[:extra])
theta = tau_star * sps_table.lookup(num_requests + extra)
if theta > best_theta:
best_theta, best_extra = theta, extra
return best_extra
def schedule_verify_lens_topk_vanilla(
*,
survival_probs: torch.Tensor,
budget: int,
cfg: DSparkScheduleConfig,
) -> torch.Tensor:
cfg.validate()
num_requests, _gamma = survival_probs.shape
max_len = cfg.resolved_max_verify_len()
device = survival_probs.device
valid_rows = (survival_probs >= cfg.survival_eps).tolist()
survival_rows = survival_probs.to(torch.float64).tolist()
candidates: list[tuple[float, int, int]] = []
for request in range(num_requests):
for position in range(min(max_len, len(valid_rows[request]))):
if valid_rows[request][position]:
candidates.append((survival_rows[request][position], position, request))
candidates.sort(key=lambda candidate: (-candidate[0], candidate[1], candidate[2]))
selected_extra = [0] * num_requests
for _survival, _position, request in candidates[: max(int(budget), 0)]:
selected_extra[request] += 1
lower_bound = max(cfg.min_verify_len, 1)
verify_lens = [
min(max(cfg.min_verify_len + extra, lower_bound), max_len)
for extra in selected_extra
]
return torch.tensor(verify_lens, dtype=torch.int32, device=device)
_TOPK_IMPLS = (
schedule_verify_lens_topk_from_survival,
schedule_verify_lens_topk_vanilla,
)
def _for_each_impl(test_method):
@functools.wraps(test_method)
def wrapper(self):
for impl in _TOPK_IMPLS:
with self.subTest(impl=impl.__name__):
test_method(self, impl)
return wrapper
class TestComputeVerifyTokenBudget(CustomTestCase):
def test_budget_argmax_matches_bruteforce_scan_across_sps_cliffs(self):
torch.manual_seed(1)
cfg = DSparkScheduleConfig(gamma=7)
table = _cliff_table()
for _ in range(20):
confidence = torch.rand(3, 7, dtype=torch.float32) * 0.5 + 0.45
survival = _survival_from_confidence(confidence)
expected = _bruteforce_budget(
history_survival_probs=survival, sps_table=table, cfg=cfg
)
actual = compute_verify_token_budget(
history_survival_probs=survival, sps_table=table, cfg=cfg
).budget
self.assertEqual(actual, expected)
def test_budget_pool_is_independent_of_min_verify_len(self):
survival = torch.tensor([[0.9, 0.8, 0.7, 0.6]], dtype=torch.float32)
cfg_no_min = DSparkScheduleConfig(gamma=4, min_verify_len=0)
cfg_min2 = DSparkScheduleConfig(gamma=4, min_verify_len=2)
table = _flat_table()
budget_no_min = compute_verify_token_budget(
history_survival_probs=survival, sps_table=table, cfg=cfg_no_min
).budget
budget_min2 = compute_verify_token_budget(
history_survival_probs=survival, sps_table=table, cfg=cfg_min2
).budget
# The budget candidate pool spans all positions; the per-request floor
# is enforced later at verify-lens scheduling, not at budget time.
self.assertEqual(budget_no_min, 4)
self.assertEqual(budget_min2, 4)
def test_budget_drops_candidates_below_survival_eps(self):
survival = torch.tensor([[0.9, 1e-9, 1e-12]], dtype=torch.float32)
cfg = DSparkScheduleConfig(gamma=3, survival_eps=1e-6)
table = _flat_table()
budget = compute_verify_token_budget(
history_survival_probs=survival, sps_table=table, cfg=cfg
).budget
self.assertLessEqual(budget, 1)
def test_decision_predicted_step_matches_additive_table_at_budget(self):
survival = torch.tensor([[0.9, 0.8, 0.7, 0.6]], dtype=torch.float32)
cfg = DSparkScheduleConfig(gamma=4)
table = _additive_table()
decision = compute_verify_token_budget(
history_survival_probs=survival, sps_table=table, cfg=cfg
)
# Reference via the independent scalar interpolation path (step_time),
# not the tensor helper the implementation itself uses.
self.assertAlmostEqual(
decision.predicted_step_seconds,
table.step_time(num_reqs=1, budget=int(decision.budget)),
places=5,
)
self.assertGreater(decision.predicted_theta, 0.0)
def test_decision_predicted_step_is_inverse_sps_for_diagonal_table(self):
survival = torch.tensor([[0.9, 0.8, 0.7, 0.6]], dtype=torch.float32)
cfg = DSparkScheduleConfig(gamma=4)
table = _cliff_table()
decision = compute_verify_token_budget(
history_survival_probs=survival, sps_table=table, cfg=cfg
)
expected_sps = table.lookup(1 + decision.budget)
self.assertIsNotNone(decision.predicted_step_seconds)
self.assertAlmostEqual(
decision.predicted_step_seconds, 1.0 / expected_sps, places=9
)
self.assertGreater(decision.predicted_theta, 0.0)
def _make_budget_planner() -> HostConfidenceBudgetPlanner:
return HostConfidenceBudgetPlanner(
sps_table=_flat_table(),
cfg=DSparkScheduleConfig(gamma=4),
model_runner=None,
)
class TestBudgetDecisionLifecycle(CustomTestCase):
def test_take_last_decision_is_consume_once(self):
planner = _make_budget_planner()
planner.last_decision = VerifyBudgetDecision(
budget=3, predicted_step_seconds=0.01, predicted_theta=100.0
)
first = planner.take_last_decision()
self.assertEqual(first.budget, 3)
self.assertIsNone(planner.take_last_decision())
def test_note_non_decode_step_clears_decision(self):
planner = _make_budget_planner()
planner.last_decision = VerifyBudgetDecision(budget=1)
planner.note_non_decode_step()
self.assertIsNone(planner.take_last_decision())
class TestScheduleVerifyLensTopk(CustomTestCase):
@_for_each_impl
def test_topk_does_not_exceed_budget(self, impl):
torch.manual_seed(2)
survival = _survival_from_confidence(torch.rand(5, 7) * 0.4 + 0.55)
cfg = DSparkScheduleConfig(gamma=7)
floor = max(cfg.min_verify_len, 1)
for budget in (0, 1, 5, 12, 100):
verify_lens = impl(survival_probs=survival, budget=budget, cfg=cfg)
total_extra = int((verify_lens.to(torch.int64) - floor).sum().item())
self.assertLessEqual(total_extra, budget)
self.assertGreaterEqual(int(verify_lens.min().item()), 1)
@_for_each_impl
def test_total_equals_anchors_plus_lens(self, impl):
survival = torch.tensor(
[[0.90, 0.80, 0.30, 0.20], [0.85, 0.70, 0.25, 0.15]],
dtype=torch.float32,
)
num_requests, max_len, budget = 2, 4, 2
cfg = DSparkScheduleConfig(gamma=max_len, min_verify_len=1)
verify_lens = impl(survival_probs=survival, budget=budget, cfg=cfg)
actual_total = num_requests + int(verify_lens.to(torch.int64).sum().item())
admitted = budget
expected_total = num_requests + num_requests * cfg.min_verify_len + admitted
self.assertEqual(actual_total, expected_total)
@_for_each_impl
def test_higher_confidence_admitted_first(self, impl):
survival = torch.tensor(
[[0.99, 0.98, 0.97, 0.96], [0.40, 0.30, 0.20, 0.10]],
dtype=torch.float32,
)
cfg = DSparkScheduleConfig(gamma=4)
verify_lens = impl(survival_probs=survival, budget=2, cfg=cfg)
extra = verify_lens.to(torch.int64) - cfg.min_verify_len
self.assertEqual(int(extra[0].item()), 2)
self.assertEqual(int(extra[1].item()), 0)
@_for_each_impl
def test_min_and_max_enter_the_budget(self, impl):
survival = torch.tensor([[0.99, 0.99, 0.99, 0.99, 0.99]], dtype=torch.float32)
cfg = DSparkScheduleConfig(gamma=5, min_verify_len=1, max_verify_len=3)
verify_lens = impl(survival_probs=survival, budget=100, cfg=cfg)
self.assertGreaterEqual(int(verify_lens.min().item()), 1)
self.assertLessEqual(int(verify_lens.max().item()), 3)
@_for_each_impl
def test_large_budget_selects_all_candidates(self, impl):
survival = torch.tensor([[0.9, 0.8, 0.7]], dtype=torch.float32)
cfg = DSparkScheduleConfig(gamma=3)
verify_lens = impl(survival_probs=survival, budget=1000, cfg=cfg)
# anchor (min_verify_len=1) + all 3 admitted drafts
self.assertEqual(int(verify_lens[0].item()), 4)
@_for_each_impl
def test_tie_break_is_value_independent(self, impl):
survival = torch.tensor([[0.8, 0.8, 0.8], [0.8, 0.8, 0.8]], dtype=torch.float32)
cfg = DSparkScheduleConfig(gamma=3)
floor = max(cfg.min_verify_len, 1)
verify_lens = impl(survival_probs=survival, budget=3, cfg=cfg)
total_extra = int((verify_lens.to(torch.int64) - floor).sum().item())
self.assertEqual(total_extra, 3)
class TestVerifyLenAnchorContract(CustomTestCase):
@_for_each_impl
def test_explicit_zero_min_still_clamped_to_anchor(self, impl):
survival = _survival_from_confidence(
torch.tensor([[0.9, 0.8, 0.7], [0.6, 0.5, 0.4]], dtype=torch.float32)
)
cfg = DSparkScheduleConfig(gamma=3, min_verify_len=0)
verify_lens = impl(survival_probs=survival, budget=0, cfg=cfg)
self.assertGreaterEqual(int(verify_lens.min().item()), 1)
self.assertTrue(
torch.equal(verify_lens, torch.tensor([1, 1], dtype=torch.int32))
)
def test_non_flat_table_small_budget_feeds_ragged_layout(self):
table = SpsCostTable(
sample_batch_tokens=[2, 3],
sample_steps_per_sec=[1.0, 0.1],
max_batch_tokens=64,
)
cfg = DSparkScheduleConfig(gamma=3)
survival = _survival_from_confidence(
torch.tensor([[0.90, 0.80, 0.70], [0.85, 0.60, 0.40]], dtype=torch.float32)
)
budget = compute_verify_token_budget(
history_survival_probs=survival, sps_table=table, cfg=cfg
).budget
self.assertEqual(budget, 0)
verify_lens = schedule_verify_lens_topk_from_survival(
survival_probs=survival, budget=budget, cfg=cfg
)
self.assertGreaterEqual(int(verify_lens.min().item()), 1)
verify_lens_cpu = verify_lens.to(torch.int64).tolist()
layout = RaggedVerifyLayout.from_verify_lens(
verify_lens_cpu=verify_lens_cpu,
device=torch.device("cpu"),
grid=[sum(verify_lens_cpu)],
)
self.assertEqual(layout.verify_lens_cpu, verify_lens_cpu)
class TestNonAnticipating(CustomTestCase):
@_for_each_impl
def test_lens_topk_non_anticipating_under_future_perturbation(self, impl):
base = torch.tensor(
[
[0.95, 0.90, 0.80, 0.40],
[0.92, 0.70, 0.30, 0.10],
[0.99, 0.98, 0.50, 0.05],
],
dtype=torch.float32,
)
cfg = DSparkScheduleConfig(gamma=4)
budget = 5
baseline = impl(survival_probs=base, budget=budget, cfg=cfg)
request, cut = 1, 2
for delta in (-0.05, -0.2, 0.05, 0.0):
perturbed = base.clone()
future = perturbed[request, cut:]
perturbed[request, cut:] = torch.clamp(
torch.minimum(future + delta, base[request, cut - 1]), min=0.0
)
verify_lens = impl(survival_probs=perturbed, budget=budget, cfg=cfg)
admitted_prefix_unchanged = min(int(baseline[request].item()), cut) == min(
int(verify_lens[request].item()), cut
)
self.assertTrue(
admitted_prefix_unchanged,
msg=f"prefix admission changed under future perturbation delta={delta}",
)
@_for_each_impl
def test_other_requests_unaffected_by_one_request_future(self, impl):
base = torch.tensor(
[[0.95, 0.90, 0.20], [0.93, 0.88, 0.15]], dtype=torch.float32
)
cfg = DSparkScheduleConfig(gamma=3)
budget = 2
baseline = impl(survival_probs=base, budget=budget, cfg=cfg)
perturbed = base.clone()
perturbed[0, 2] = 0.01
verify_lens = impl(survival_probs=perturbed, budget=budget, cfg=cfg)
self.assertEqual(int(baseline[1].item()), int(verify_lens[1].item()))
class TestVanillaMatchesReference(CustomTestCase):
def test_random_inputs_match_reference(self):
torch.manual_seed(20260630)
num_trials = 4000
for trial in range(num_trials):
num_requests = int(torch.randint(1, 6, ()).item())
gamma = int(torch.randint(1, 9, ()).item())
dtype = torch.float32 if trial % 2 == 0 else torch.float64
confidence = torch.rand(num_requests, gamma, dtype=dtype)
if trial % 3 == 0:
confidence = (confidence * 4).round() / 4
if trial % 7 == 0:
confidence = torch.ones(num_requests, gamma, dtype=dtype)
survival = torch.cumprod(confidence, dim=1)
min_verify_len = int(torch.randint(0, gamma + 1, ()).item())
if torch.rand(()).item() < 0.5:
max_verify_len = 0
else:
max_verify_len = int(
torch.randint(min_verify_len, gamma + 1, ()).item()
)
survival_eps = float(
[1e-6, 1e-3, 0.1, 0.5][int(torch.randint(0, 4, ()).item())]
)
budget = int(torch.randint(0, num_requests * gamma + 3, ()).item())
cfg = DSparkScheduleConfig(
gamma=gamma,
min_verify_len=min_verify_len,
max_verify_len=max_verify_len,
survival_eps=survival_eps,
)
reference = schedule_verify_lens_topk_from_survival(
survival_probs=survival, budget=budget, cfg=cfg
)
vanilla = schedule_verify_lens_topk_vanilla(
survival_probs=survival, budget=budget, cfg=cfg
)
self.assertTrue(
torch.equal(reference, vanilla),
msg=(
f"mismatch on trial {trial}: budget={budget} "
f"min={min_verify_len} max={max_verify_len} eps={survival_eps} "
f"survival={survival.tolist()} "
f"reference={reference.tolist()} vanilla={vanilla.tolist()}"
),
)
class TestDSparkScheduleConfig(CustomTestCase):
def test_validate_rejects_min_greater_than_max(self):
with self.assertRaises(ValueError):
DSparkScheduleConfig(gamma=4, min_verify_len=3, max_verify_len=2).validate()
def test_validate_rejects_max_greater_than_gamma_plus_one(self):
with self.assertRaises(ValueError):
DSparkScheduleConfig(gamma=4, max_verify_len=6).validate()
def test_zero_max_resolves_to_gamma_plus_one(self):
cfg = DSparkScheduleConfig(gamma=7)
self.assertEqual(cfg.resolved_max_verify_len(), 8)
class TestGraphTierFillBudget(CustomTestCase):
def test_floor_scales_with_min_verify_len(self):
"""The subtracted floor is bs * max(min_verify_len, 1)."""
self.assertEqual(
graph_tier_fill_budget(
graph_num_tokens=60, bs=10, verify_num_draft_tokens=6, min_verify_len=2
),
60 - 20,
)
def test_feeding_budget_fills_topk_total_to_tier(self):
"""Feeding the fill budget to the top-k lifts the total to min(tier, bs*K)."""
bs = 6
cfg = DSparkScheduleConfig(gamma=6, min_verify_len=1)
cap = cfg.resolved_max_verify_len()
survival = torch.full((bs, cap), 0.99, dtype=torch.float32)
for graph_num_tokens in (bs, bs * cap // 2, bs * cap, bs * cap + 12):
budget = graph_tier_fill_budget(
graph_num_tokens=graph_num_tokens,
bs=bs,
verify_num_draft_tokens=cap,
min_verify_len=cfg.min_verify_len,
)
verify_lens = schedule_verify_lens_topk_from_survival(
survival_probs=survival, budget=budget, cfg=cfg
)
total = int(verify_lens.to(torch.int64).sum().item())
self.assertEqual(total, min(graph_num_tokens, bs * cap))
class _FakeRaggedRunner(types.SimpleNamespace):
pass
def _fake_model_runner(capture_num_tokens, max_bs):
runner = _FakeRaggedRunner(
ragged_verify_mode=True,
capture_num_tokens=capture_num_tokens,
max_bs=max_bs,
)
return types.SimpleNamespace(decode_cuda_graph_runner=runner)
class TestBudgetTierSelection(CustomTestCase):
def test_floor_uses_tier_hint_capped_at_uniform_window(self):
from sglang.srt.speculative.dspark_components.dspark_planner import (
verify_layout_graph_num_tokens_floor,
)
from sglang.srt.speculative.ragged_verify import RaggedVerifyMode
model_runner = _fake_model_runner([8, 16, 1024], max_bs=128)
floor = verify_layout_graph_num_tokens_floor(
num_reqs=100,
ragged_verify_mode=RaggedVerifyMode.COMPACT,
verify_num_draft_tokens=8,
model_runner=model_runner,
tier_num_tokens=150,
)
self.assertEqual(floor, 150)
capped = verify_layout_graph_num_tokens_floor(
num_reqs=10,
ragged_verify_mode=RaggedVerifyMode.COMPACT,
verify_num_draft_tokens=8,
model_runner=model_runner,
tier_num_tokens=150,
)
self.assertEqual(capped, 80)
pinned = verify_layout_graph_num_tokens_floor(
num_reqs=100,
ragged_verify_mode=RaggedVerifyMode.COMPACT,
verify_num_draft_tokens=8,
model_runner=model_runner,
)
self.assertEqual(pinned, 800)
def test_exceeds_gate_checks_slots_and_tier(self):
from sglang.srt.speculative.dspark_components.dspark_planner import (
ragged_layout_exceeds_captured_grid,
)
model_runner = _fake_model_runner([8, 16, 1024], max_bs=128)
self.assertTrue(
ragged_layout_exceeds_captured_grid(
num_reqs=129,
verify_num_draft_tokens=8,
model_runner=model_runner,
tier_tokens_hint=200,
)
)
self.assertFalse(
ragged_layout_exceeds_captured_grid(
num_reqs=128,
verify_num_draft_tokens=8,
model_runner=model_runner,
tier_tokens_hint=512,
)
)
self.assertFalse(
ragged_layout_exceeds_captured_grid(
num_reqs=128,
verify_num_draft_tokens=8,
model_runner=model_runner,
)
)
self.assertTrue(
ragged_layout_exceeds_captured_grid(
num_reqs=128,
verify_num_draft_tokens=9,
model_runner=model_runner,
)
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,278 @@
import unittest
from sglang.benchmark.dspark_sps_profiler import (
LoadInfo,
ServerContext,
SpsRow,
build_request_count_sweep,
build_table_from_summaries,
count_aligned_steps,
postprocess_round,
resolve_cuda_graph_max_bs,
round_summary_dict,
validate_sweep_against_server,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=15, suite="base-a-test-cpu")
def make_load_info() -> LoadInfo:
return LoadInfo(
num_requests=4, max_new_tokens=1200, wall_seconds=1.0, reached_target=True
)
def make_rows(
*,
num_rows: int = 30,
num_running_reqs: int = 4,
num_verify_tokens: int = 32,
step_time: float = 0.01,
first_forward_ct: int = 0,
) -> list[SpsRow]:
return [
SpsRow(
forward_ct=first_forward_ct + index,
num_running_reqs=num_running_reqs,
num_verify_tokens=num_verify_tokens,
step_time=step_time,
)
for index in range(num_rows)
]
def make_context(**overrides) -> ServerContext:
values = dict(
base_url="http://localhost:30000",
tokenizer_path="dummy",
tp_size=4,
dp_size=1,
verify_num_draft_tokens=8,
simulate_acc_len=1.0,
cuda_graph_max_bs=128,
skip_max_running_requests_threshold=float("inf"),
skip_token_capacity_threshold=float("inf"),
)
values.update(overrides)
return ServerContext(**values)
class TestPostprocessRound(CustomTestCase):
def test_single_rank_round_builds_probe_from_median_step_time(self):
outcome = postprocess_round(
rank_rows=[make_rows(step_time=0.01)],
batch_size_per_rank=4,
dp_size=1,
verify_num_draft_tokens=8,
min_steady_steps=16,
load_info=make_load_info(),
)
self.assertEqual(outcome.batch_tokens, 32)
self.assertAlmostEqual(outcome.steps_per_sec, 100.0)
self.assertEqual(outcome.match_fraction, 1.0)
def test_round_warmup_steps_are_dropped_from_timing(self):
slow_head = make_rows(num_rows=8, step_time=0.5, first_forward_ct=0)
steady_tail = make_rows(num_rows=20, step_time=0.01, first_forward_ct=8)
outcome = postprocess_round(
rank_rows=[slow_head + steady_tail],
batch_size_per_rank=4,
dp_size=1,
verify_num_draft_tokens=8,
min_steady_steps=16,
load_info=make_load_info(),
)
self.assertAlmostEqual(outcome.steps_per_sec, 100.0)
def test_off_target_batch_rows_are_filtered_out(self):
ramp = make_rows(num_rows=10, num_running_reqs=2, num_verify_tokens=16)
steady = make_rows(num_rows=30, first_forward_ct=10, step_time=0.02)
outcome = postprocess_round(
rank_rows=[ramp + steady],
batch_size_per_rank=4,
dp_size=1,
verify_num_draft_tokens=8,
min_steady_steps=16,
load_info=make_load_info(),
)
self.assertAlmostEqual(outcome.steps_per_sec, 50.0)
self.assertAlmostEqual(outcome.match_fraction, 1.0)
def test_mid_round_instability_raises(self):
head = make_rows(num_rows=15)
gap = make_rows(
num_rows=40, num_running_reqs=3, num_verify_tokens=24, first_forward_ct=15
)
tail = make_rows(num_rows=15, first_forward_ct=55)
with self.assertRaisesRegex(RuntimeError, "unstable mid-round"):
postprocess_round(
rank_rows=[head + gap + tail],
batch_size_per_rank=4,
dp_size=1,
verify_num_draft_tokens=8,
min_steady_steps=16,
load_info=make_load_info(),
)
def test_round_that_never_stabilizes_raises(self):
rows = make_rows(num_rows=50, num_running_reqs=3, num_verify_tokens=24)
rows += make_rows(num_rows=2, first_forward_ct=50)
with self.assertRaisesRegex(RuntimeError, "never stabilized"):
postprocess_round(
rank_rows=[rows],
batch_size_per_rank=4,
dp_size=1,
verify_num_draft_tokens=8,
min_steady_steps=16,
load_info=make_load_info(),
)
class TestPostprocessRoundCrossRank(CustomTestCase):
def test_two_uniform_ranks_average_their_step_times(self):
outcome = postprocess_round(
rank_rows=[make_rows(step_time=0.01), make_rows(step_time=0.03)],
batch_size_per_rank=4,
dp_size=2,
verify_num_draft_tokens=8,
min_steady_steps=16,
load_info=make_load_info(),
)
self.assertEqual(outcome.batch_size_per_rank, 4)
self.assertEqual(outcome.batch_tokens, 32)
self.assertAlmostEqual(outcome.steps_per_sec, 50.0)
self.assertEqual(len(outcome.per_rank_median_step_time), 2)
self.assertAlmostEqual(outcome.per_rank_median_step_time[0], 0.01)
self.assertAlmostEqual(outcome.per_rank_median_step_time[1], 0.03)
def test_rank_with_no_new_records_raises(self):
with self.assertRaisesRegex(RuntimeError, "no new decode-step records"):
postprocess_round(
rank_rows=[make_rows(), []],
batch_size_per_rank=4,
dp_size=2,
verify_num_draft_tokens=8,
min_steady_steps=16,
load_info=make_load_info(),
)
def test_disjoint_forward_ct_ranges_raise(self):
with self.assertRaisesRegex(RuntimeError, "no common forward_ct"):
postprocess_round(
rank_rows=[
make_rows(first_forward_ct=0),
make_rows(first_forward_ct=1000),
],
batch_size_per_rank=4,
dp_size=2,
verify_num_draft_tokens=8,
min_steady_steps=16,
load_info=make_load_info(),
)
def test_rank_below_expected_verify_tokens_raises(self):
# A rank reporting fewer verify tokens than bs_per_rank * K is not
# running the uniform static verify; above-expected counts are
# tolerated (the recorded count is the replayed graph tier).
with self.assertRaisesRegex(RuntimeError, "num_verify_tokens"):
postprocess_round(
rank_rows=[make_rows(), make_rows(num_verify_tokens=24)],
batch_size_per_rank=4,
dp_size=2,
verify_num_draft_tokens=8,
min_steady_steps=16,
load_info=make_load_info(),
)
def test_rank_count_mismatch_raises(self):
with self.assertRaisesRegex(RuntimeError, "DP ranks"):
postprocess_round(
rank_rows=[make_rows()],
batch_size_per_rank=4,
dp_size=2,
verify_num_draft_tokens=8,
min_steady_steps=16,
load_info=make_load_info(),
)
class TestTableAssembly(CustomTestCase):
def test_repeats_take_the_median_per_batch_tokens(self):
rounds = [
postprocess_round(
rank_rows=[make_rows(step_time=step_time)],
batch_size_per_rank=4,
dp_size=1,
verify_num_draft_tokens=8,
min_steady_steps=16,
load_info=make_load_info(),
)
for step_time in (0.01, 0.02, 0.04)
]
table = build_table_from_summaries(
summaries=[
round_summary_dict(outcome=outcome, repeat=repeat)
for repeat, outcome in enumerate(rounds)
],
max_batch_tokens=None,
offdiag=False,
)
self.assertEqual(table.sample_batch_tokens, [32])
self.assertAlmostEqual(table.sample_steps_per_sec[0], 50.0)
class TestSweepHelpers(CustomTestCase):
def test_request_count_sweep_tapers_and_hits_the_max(self):
sweep = build_request_count_sweep(100)
self.assertEqual(sweep[:4], [1, 2, 4, 8])
self.assertEqual(sweep[-1], 100)
self.assertIn(64, sweep)
def test_sweep_beyond_captured_cuda_graphs_raises(self):
with self.assertRaisesRegex(ValueError, "cuda graphs"):
validate_sweep_against_server(
context=make_context(cuda_graph_max_bs=64),
batch_sizes=[8, 128],
)
def test_sweep_within_captured_cuda_graphs_passes(self):
validate_sweep_against_server(
context=make_context(cuda_graph_max_bs=64, dp_size=2),
batch_sizes=[8, 64],
)
def test_resolve_cuda_graph_max_bs_prefers_captured_list(self):
internal_state = {
"cuda_graph_config": {"decode": {"bs": [1, 2, 160], "max_bs": 128}}
}
self.assertEqual(resolve_cuda_graph_max_bs(internal_state=internal_state), 160)
def test_resolve_cuda_graph_max_bs_handles_missing_config(self):
self.assertIsNone(resolve_cuda_graph_max_bs(internal_state={}))
class TestCountAlignedSteps(CustomTestCase):
def test_off_target_steps_are_not_counted(self):
rows = make_rows(num_rows=10, num_running_reqs=3)
self.assertEqual(
count_aligned_steps(rank_rows=[rows], batch_size_per_rank=4), 0
)
class TestMinSteadySteps(CustomTestCase):
def test_min_steady_steps_rejects_thin_probes(self):
with self.assertRaisesRegex(RuntimeError, "never stabilized"):
postprocess_round(
rank_rows=[make_rows(num_rows=20)],
batch_size_per_rank=4,
dp_size=1,
verify_num_draft_tokens=8,
min_steady_steps=32,
load_info=make_load_info(),
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,200 @@
import tempfile
import unittest
from pathlib import Path
from types import SimpleNamespace
from sglang.srt.speculative.dspark_components.dspark_sps import (
SpsAdditiveCostTable,
SpsCostTable,
build_uninitialized_sps_table,
is_uninitialized_sps_table,
load_sps_table_from_path,
profile_sps_table,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
def _make_table() -> SpsCostTable:
return SpsCostTable(
sample_batch_tokens=[8, 16, 32, 64],
sample_steps_per_sec=[1000.0, 950.0, 500.0, 480.0],
max_batch_tokens=128,
)
class TestSpsCostTableInvariants(CustomTestCase):
def test_rejects_non_increasing_batch_tokens(self):
with self.assertRaises(ValueError):
SpsCostTable(
sample_batch_tokens=[8, 8, 16],
sample_steps_per_sec=[1.0, 2.0, 3.0],
max_batch_tokens=16,
)
def test_rejects_unsorted_batch_tokens(self):
with self.assertRaises(ValueError):
SpsCostTable(
sample_batch_tokens=[16, 8],
sample_steps_per_sec=[1.0, 2.0],
max_batch_tokens=16,
)
def test_rejects_length_mismatch(self):
with self.assertRaises(ValueError):
SpsCostTable(
sample_batch_tokens=[8, 16],
sample_steps_per_sec=[1.0],
max_batch_tokens=16,
)
def test_rejects_empty_table(self):
with self.assertRaises(ValueError):
SpsCostTable(
sample_batch_tokens=[],
sample_steps_per_sec=[],
max_batch_tokens=0,
)
def test_rejects_max_below_largest_probe(self):
with self.assertRaises(ValueError):
SpsCostTable(
sample_batch_tokens=[8, 16],
sample_steps_per_sec=[1.0, 2.0],
max_batch_tokens=15,
)
class TestSpsCostTableLookup(CustomTestCase):
def test_lookup_exact_probe_returns_that_sps(self):
table = _make_table()
self.assertEqual(table.lookup(8), 1000.0)
self.assertEqual(table.lookup(16), 950.0)
self.assertEqual(table.lookup(32), 500.0)
self.assertEqual(table.lookup(64), 480.0)
def test_lookup_floors_to_lower_captured_probe(self):
table = _make_table()
self.assertEqual(table.lookup(31), 950.0)
self.assertEqual(table.lookup(63), 500.0)
def test_lookup_below_first_probe_clamps_to_first(self):
table = _make_table()
self.assertEqual(table.lookup(1), 1000.0)
self.assertEqual(table.lookup(7), 1000.0)
def test_lookup_above_last_probe_clamps_to_last(self):
table = _make_table()
self.assertEqual(table.lookup(65), 480.0)
self.assertEqual(table.lookup(10_000), 480.0)
class TestLoadSpsTableFromPath(CustomTestCase):
def test_load_from_path_round_trips_table_and_lookup(self):
table = _make_table()
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "sps.json"
path.write_text(table.to_json(), encoding="utf-8")
loaded = load_sps_table_from_path(str(path))
self.assertEqual(loaded.sample_batch_tokens, table.sample_batch_tokens)
self.assertEqual(loaded.sample_steps_per_sec, table.sample_steps_per_sec)
self.assertEqual(loaded.max_batch_tokens, table.max_batch_tokens)
for batch_tokens in (1, 8, 31, 64, 200):
self.assertEqual(loaded.lookup(batch_tokens), table.lookup(batch_tokens))
class TestFlatTableLookupIsConstant(CustomTestCase):
def test_flat_table_lookup_is_one_for_any_batch(self):
flat = SpsCostTable(
sample_batch_tokens=[1],
sample_steps_per_sec=[1.0],
max_batch_tokens=4096,
)
for batch_tokens in (0, 1, 2, 17, 256, 100_000):
self.assertEqual(flat.lookup(batch_tokens), 1.0)
class TestProfileSpsTable(CustomTestCase):
def test_profile_sorts_out_of_order_probes(self):
table = profile_sps_table(
probes=[(32, 500.0), (8, 1000.0), (16, 950.0)],
)
self.assertEqual(table.sample_batch_tokens, [8, 16, 32])
self.assertEqual(table.sample_steps_per_sec, [1000.0, 950.0, 500.0])
def test_profile_rejects_duplicate_batch_tokens(self):
with self.assertRaises(ValueError):
profile_sps_table(probes=[(8, 1000.0), (8, 900.0)])
def test_profile_rejects_empty_probes(self):
with self.assertRaises(ValueError):
profile_sps_table(probes=[])
def test_profile_max_batch_tokens_defaults_to_largest_probe(self):
table = profile_sps_table(probes=[(8, 1000.0), (64, 480.0), (16, 950.0)])
self.assertEqual(table.max_batch_tokens, 64)
def test_profile_honors_explicit_max_batch_tokens(self):
table = profile_sps_table(
probes=[(8, 1000.0), (16, 950.0)], max_batch_tokens=256
)
self.assertEqual(table.max_batch_tokens, 256)
def _build_sps_cost_table_for(*, sps_table_path):
from sglang.srt.speculative.dspark_components.dspark_planner import (
build_sps_cost_table,
)
server_args = SimpleNamespace(
speculative_dspark_sps_table_path=sps_table_path,
max_running_requests=4,
)
return build_sps_cost_table(server_args=server_args, verify_num_draft_tokens=5)
class TestBuildSpsCostTableContract(CustomTestCase):
def test_unset_table_path_returns_flat_table(self):
for sps_table_path in (None, ""):
table = _build_sps_cost_table_for(sps_table_path=sps_table_path)
self.assertEqual(table.sample_batch_tokens, [1])
self.assertEqual(table.sample_steps_per_sec, [1.0])
self.assertEqual(table.max_batch_tokens, 20)
def test_real_path_loads_table(self):
table = _make_table()
with tempfile.TemporaryDirectory() as tmp:
path = Path(tmp) / "sps.json"
path.write_text(table.to_json(), encoding="utf-8")
loaded = _build_sps_cost_table_for(sps_table_path=str(path))
self.assertEqual(loaded.sample_batch_tokens, table.sample_batch_tokens)
self.assertEqual(loaded.sample_steps_per_sec, table.sample_steps_per_sec)
self.assertEqual(loaded.max_batch_tokens, table.max_batch_tokens)
class TestIsUninitializedSpsTable(CustomTestCase):
def test_additive_table_is_never_uninitialized(self):
table = SpsAdditiveCostTable(
bias_seconds=0.1,
bs_probes=[128, 192, 256],
alpha_seconds=[0.0, 0.008, 0.016],
m_probes=[384, 512, 1024],
theta_seconds=[0.0, 0.02, 0.1],
)
self.assertFalse(is_uninitialized_sps_table(table))
def test_placeholder_diagonal_table_is_uninitialized(self):
self.assertTrue(
is_uninitialized_sps_table(
build_uninitialized_sps_table(max_batch_tokens=128)
)
)
def test_real_diagonal_table_is_initialized(self):
self.assertFalse(is_uninitialized_sps_table(_make_table()))
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,142 @@
import tempfile
import unittest
from pathlib import Path
import torch
from sglang.benchmark.dspark_sts_fit import (
default_temperature_grid,
expected_calibration_error,
fit_sts_temperatures,
)
from sglang.srt.models.dspark import DSparkConfidenceHead
from sglang.srt.speculative.dspark_components.dspark_sts import (
DSparkStsCalibration,
StsDataRecorder,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
class TestApplySts(CustomTestCase):
def test_default_buffer_is_identity_sigmoid(self):
head = DSparkConfidenceHead(hidden_size=8, markov_rank=4, with_markov=False)
confidence_raw = torch.randn(3, 5) * 9.0
out = head.apply_sts(confidence_raw)
self.assertTrue(torch.equal(out, torch.sigmoid(confidence_raw.float())))
def test_per_position_temperature_scales_each_column(self):
head = DSparkConfidenceHead(hidden_size=8, markov_rank=4, with_markov=False)
head.sts_temperatures = torch.tensor([0.5, 1.0, 2.0])
confidence_raw = torch.full((2, 3), 2.0)
out = head.apply_sts(confidence_raw)
# Hand-computed sigmoid(2.0 / T) per column: identical raw logits with
# distinct per-column T catch wrong-axis broadcasts, and dividing (not
# multiplying) by T is what separates 0.982 from 0.731 in column 0.
expected_row = [0.98201379, 0.88079708, 0.73105858]
for row in out.tolist():
for got, want in zip(row, expected_row):
self.assertAlmostEqual(got, want, places=6)
def test_apply_sts_stashes_raw_logit(self):
head = DSparkConfidenceHead(hidden_size=8, markov_rank=4, with_markov=False)
confidence_raw = torch.randn(2, 5)
head.apply_sts(confidence_raw)
self.assertIs(head._last_confidence_raw, confidence_raw)
class TestDSparkStsCalibration(CustomTestCase):
def test_json_round_trip_preserves_fields(self):
calibration = DSparkStsCalibration(
temperatures=[1.5, 2.0, 0.5],
dataset="shards.*.pt",
num_samples=1234,
ece_before=[0.3, 0.2, 0.1],
ece_after=[0.02, 0.01, 0.03],
)
restored = DSparkStsCalibration.from_json(calibration.to_json())
self.assertEqual(restored.temperatures, calibration.temperatures)
self.assertEqual(restored.dataset, calibration.dataset)
self.assertEqual(restored.num_samples, calibration.num_samples)
self.assertEqual(restored.ece_before, calibration.ece_before)
self.assertEqual(restored.ece_after, calibration.ece_after)
def test_rejects_empty_temperatures(self):
with self.assertRaises(ValueError):
DSparkStsCalibration(temperatures=[])
def test_rejects_non_positive_temperature(self):
with self.assertRaises(ValueError):
DSparkStsCalibration(temperatures=[1.0, 0.0, 2.0])
with self.assertRaises(ValueError):
DSparkStsCalibration(temperatures=[1.0, -0.5])
class TestExpectedCalibrationError(CustomTestCase):
def test_perfectly_calibrated_probs_have_low_ece(self):
torch.manual_seed(0)
probs = torch.full((20000,), 0.3, dtype=torch.float64)
targets = (torch.rand(20000) < 0.3).to(torch.float64)
ece = expected_calibration_error(probs=probs, targets=targets, num_bins=15)
self.assertLess(ece, 0.02)
def test_overconfident_probs_have_high_ece(self):
probs = torch.full((20000,), 0.95, dtype=torch.float64)
targets = torch.full((20000,), 0.3, dtype=torch.float64)
ece = expected_calibration_error(probs=probs, targets=targets, num_bins=15)
self.assertGreater(ece, 0.5)
class TestFitStsTemperatures(CustomTestCase):
def test_recovers_scale_and_reduces_ece(self):
torch.manual_seed(0)
num_samples, gamma, scale = 60000, 4, 2.5
base_logit = torch.tensor([2.0, 1.2, 0.8, 0.4])
true_logit = base_logit[None, :] + torch.randn(num_samples, gamma) * 0.5
true_prob = torch.sigmoid(true_logit)
accept = (torch.rand(num_samples, gamma) < true_prob).to(torch.float64)
prefix_mask = torch.cumprod(accept, dim=1)
overconfident_logits = true_logit * scale
result = fit_sts_temperatures(
logits=overconfident_logits,
prefix_mask=prefix_mask,
grid=default_temperature_grid(),
num_bins=15,
)
self.assertEqual(len(result["temperatures"]), gamma)
for temperature in result["temperatures"]:
self.assertGreater(temperature, scale / 1.5)
self.assertLess(temperature, scale * 1.5)
mean_before = sum(result["ece_before"]) / gamma
mean_after = sum(result["ece_after"]) / gamma
self.assertLess(mean_after, 0.25 * mean_before)
class TestStsDataRecorder(CustomTestCase):
def test_builds_prefix_mask_and_writes_shard(self):
gamma = 4
confidence_raw = torch.randn(4, gamma)
num_correct_drafts = torch.tensor([0, 2, 4, 1], dtype=torch.int32)
expected_prefix_mask = torch.tensor(
[[0, 0, 0, 0], [1, 1, 0, 0], [1, 1, 1, 1], [1, 0, 0, 0]],
dtype=torch.float32,
)
with tempfile.TemporaryDirectory() as tmp:
stem = str(Path(tmp) / "shard")
recorder = StsDataRecorder(path_stem=stem, gamma=gamma, flush_every=10)
recorder.record(
confidence_raw=confidence_raw,
num_correct_drafts=num_correct_drafts,
)
recorder.flush()
shard = torch.load(f"{stem}.0.pt")
self.assertTrue(torch.equal(shard["prefix_mask"], expected_prefix_mask))
self.assertTrue(torch.equal(shard["logits"], confidence_raw.to(torch.float32)))
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,132 @@
import unittest
import torch
from sglang.srt.speculative.ragged_verify import (
RaggedVerifyLayout,
build_ragged_target_verify_geometry,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
_DEVICE = torch.device("cpu")
_GRID = [8, 16, 24, 32, 64]
# The backend capability checks (supports_ragged_verify_graph) live in
# test_ragged_verify_backend_capability.py: importing the backend modules
# pulls GPU-only wheels, which fail to import on the CPU runners.
class TestRaggedTargetVerifyGeometry(CustomTestCase):
def test_mixed_verify_lens_geometry(self):
layout = RaggedVerifyLayout.from_verify_lens(
verify_lens_cpu=[8, 1, 3], device=_DEVICE, grid=_GRID
)
seq_lens = torch.tensor([10, 20, 30], dtype=torch.int32)
geometry = build_ragged_target_verify_geometry(seq_lens=seq_lens, layout=layout)
self.assertEqual(geometry.cache_seqlens_int32.tolist(), [18, 21, 33])
self.assertEqual(geometry.cu_seqlens_q.tolist(), [0, 8, 9, 12])
self.assertEqual(geometry.cu_seqlens_k.tolist(), [0, 18, 39, 72])
self.assertEqual(geometry.max_seq_len_q, 8)
def test_geometry_dtypes_are_int32(self):
layout = RaggedVerifyLayout.from_verify_lens(
verify_lens_cpu=[8, 1, 3], device=_DEVICE, grid=_GRID
)
seq_lens = torch.tensor([10, 20, 30], dtype=torch.int64)
geometry = build_ragged_target_verify_geometry(seq_lens=seq_lens, layout=layout)
self.assertEqual(geometry.cache_seqlens_int32.dtype, torch.int32)
self.assertEqual(geometry.cu_seqlens_q.dtype, torch.int32)
self.assertEqual(geometry.cu_seqlens_k.dtype, torch.int32)
class TestPaddedRaggedVerifyGeometry(CustomTestCase):
def test_padded_layout_grows_bs_and_fills_bucket(self):
raw = RaggedVerifyLayout.from_verify_lens(
verify_lens_cpu=[8, 1, 3],
device=_DEVICE,
grid=[8, 16, 32, 64],
graph_num_tokens_floor=24,
)
self.assertEqual(raw.graph_num_tokens, 32)
padded = raw.padded_to_bucket(padded_bs=4)
self.assertEqual(padded.bs, 4)
self.assertEqual(padded.verify_lens.tolist(), [8, 1, 3, 20])
self.assertEqual(padded.qo_indptr_device.tolist(), [0, 8, 9, 12, 32])
seq_lens = torch.tensor([10, 20, 30, 1], dtype=torch.int32)
geometry = build_ragged_target_verify_geometry(seq_lens=seq_lens, layout=padded)
self.assertEqual(geometry.cu_seqlens_q.tolist(), [0, 8, 9, 12, 32])
self.assertEqual(geometry.cache_seqlens_int32.tolist(), [18, 21, 33, 21])
self.assertEqual(int(geometry.cu_seqlens_k[-1]), 18 + 21 + 33 + 21)
def test_padded_layout_decoupled_slots_spread_slack(self):
raw = RaggedVerifyLayout.from_verify_lens(
verify_lens_cpu=[8, 1, 3],
device=_DEVICE,
grid=[8, 16, 32, 64],
graph_num_tokens_floor=24,
)
padded = raw.padded_to_bucket(padded_bs=6)
self.assertEqual(padded.bs, 6)
self.assertEqual(padded.verify_lens.tolist(), [8, 1, 3, 7, 7, 6])
self.assertEqual(int(padded.qo_indptr_device[-1]), 32)
def test_padded_layout_budget_tier_below_uniform(self):
raw = RaggedVerifyLayout.from_verify_lens(
verify_lens_cpu=[8, 1, 3],
device=_DEVICE,
grid=[8, 16, 32, 64],
)
self.assertEqual(raw.graph_num_tokens, 16)
padded = raw.padded_to_bucket(padded_bs=3)
self.assertEqual(padded.verify_lens.tolist(), [8, 1, 7])
self.assertEqual(int(padded.qo_indptr_device[-1]), 16)
def test_padded_layout_zero_len_pad_rows(self):
raw = RaggedVerifyLayout.from_verify_lens(
verify_lens_cpu=[8, 8],
device=_DEVICE,
grid=[8, 16, 32, 64],
)
self.assertEqual(raw.graph_num_tokens, 16)
padded = raw.padded_to_bucket(padded_bs=8)
self.assertEqual(padded.verify_lens.tolist(), [8, 8, 0, 0, 0, 0, 0, 0])
self.assertEqual(int(padded.qo_indptr_device[-1]), 16)
class TestCaptureVerifyLens(CustomTestCase):
def test_small_tier_one_token_rows(self):
from sglang.srt.speculative.ragged_verify import build_capture_verify_lens
lens = build_capture_verify_lens(num_tokens=8, num_slots=8, num_draft_tokens=8)
self.assertEqual(lens, [1] * 8)
def test_large_tier_spreads_within_window(self):
from sglang.srt.speculative.ragged_verify import build_capture_verify_lens
lens = build_capture_verify_lens(
num_tokens=1024, num_slots=128, num_draft_tokens=8
)
self.assertEqual(sum(lens), 1024)
self.assertEqual(lens, [8] * 128)
def test_uneven_tier_rows_stay_legal(self):
from sglang.srt.speculative.ragged_verify import build_capture_verify_lens
lens = build_capture_verify_lens(num_tokens=24, num_slots=5, num_draft_tokens=8)
self.assertEqual(sum(lens), 24)
self.assertTrue(all(1 <= v <= 8 for v in lens))
def test_rejects_overpacked_tier(self):
from sglang.srt.speculative.ragged_verify import build_capture_verify_lens
with self.assertRaises(ValueError):
build_capture_verify_lens(num_tokens=64, num_slots=4, num_draft_tokens=8)
with self.assertRaises(ValueError):
build_capture_verify_lens(num_tokens=4, num_slots=8, num_draft_tokens=8)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,43 @@
"""Backend opt-in flags for the ragged-verify graphs.
Runs in the GPU suite because importing the backend modules pulls GPU-only
wheels (sgl_kernel) at module scope, which fail to import on CPU runners.
"""
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-small")
class TestRaggedVerifyGraphCapability(CustomTestCase):
def test_base_backend_defaults_false(self):
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
self.assertFalse(AttentionBackend.supports_ragged_verify_graph)
def test_ragged_implementing_backends_declare_the_flag(self):
"""Every backend with a ragged-verify metadata path must opt in; a
dropped flag silently disables ragged graphs for that backend (the
runner falls back to eager with no other test going red)."""
from sglang.srt.layers.attention.deepseek_v4_backend import (
DeepseekV4AttnBackend,
)
from sglang.srt.layers.attention.flashattention_backend import (
FlashAttentionBackend,
)
from sglang.srt.layers.attention.trtllm_mha_backend import TRTLLMHAAttnBackend
for backend in (
TRTLLMHAAttnBackend,
DeepseekV4AttnBackend,
FlashAttentionBackend,
):
with self.subTest(backend=backend.__name__):
self.assertTrue(backend.supports_ragged_verify_graph)
if __name__ == "__main__":
unittest.main()
@@ -93,6 +93,8 @@ def _make_result(num_draft_tokens, accept_lens, flat_tokens):
speculative_num_draft_tokens=num_draft_tokens,
num_correct_drafts=None,
num_correct_drafts_per_req_cpu=None,
block_accept_lens=None,
cap_lens=None,
)
@@ -121,6 +121,7 @@ def _make_model_runner(
spec.is_eagle.return_value = False
spec.is_standalone.return_value = False
spec.is_dflash.return_value = False
spec.is_dflash_family.return_value = False
spec.is_none.return_value = True
mr.spec_algorithm = spec
@@ -30,6 +30,7 @@ def _make_info(batch_size=2, **overrides):
top_ks=torch.full((batch_size,), TOP_K_ALL, dtype=torch.int32),
min_ps=torch.zeros(batch_size),
is_all_greedy=False,
is_any_greedy=False,
need_top_p_sampling=False,
need_top_k_sampling=False,
need_min_p_sampling=False,