[Speculative Decoding] Add native UNO serving support (#37667)

Co-authored-by: drproduck <drproduck@MacBook-Air-2.local>
Co-authored-by: BBuf <1182563586@qq.com>
This commit is contained in:
Yash Akhauri
2026-09-03 20:08:41 +08:00
committed by GitHub
co-authored by drproduck BBuf
parent 354ed6d66b
commit 2bb25dc18b
51 changed files with 6300 additions and 74 deletions
@@ -0,0 +1,41 @@
"""Regression test for UNO's base-only LoRA routing fast path."""
import unittest
from types import SimpleNamespace
from sglang.srt.lora.lora_manager import LoRAManager
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class _InactiveSkippingBackend:
skip_inactive_lora_batches = True
def __init__(self):
self.batch_info = object()
self.prepare_called = False
def reset_batch_state(self):
self.batch_info = None
def prepare_lora_batch(self, *args, **kwargs):
self.prepare_called = True
class TestUnoInactiveLoRABatch(CustomTestCase):
def test_all_base_batch_clears_stale_routing_before_graph_metadata(self):
backend = _InactiveSkippingBackend()
manager = LoRAManager.__new__(LoRAManager)
manager.lora_backend = backend
forward_batch = SimpleNamespace(lora_ids=[None], batch_size=1)
manager.prepare_lora_batch(forward_batch)
self.assertIsNone(backend.batch_info)
self.assertFalse(backend.prepare_called)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,236 @@
"""Target-layer validation for UNO's specialized LoRA backend."""
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
import torch
from sglang.srt.layers.linear import (
ColumnParallelLinear,
ReplicatedLinear,
RowParallelLinear,
)
from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE
from sglang.srt.lora.backend.triton_backend import TritonLoRABackend
from sglang.srt.lora.backend.uno_cublas_backend import UnoCublasLoRABackend
from sglang.srt.lora.lora_manager import LoRAManager
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestUnoLoRATargets(CustomTestCase):
def setUp(self):
self.backend = UnoCublasLoRABackend.__new__(UnoCublasLoRABackend)
self.backend._pending_lora_a = None
self.backend._use_cublas_lora_b = False
@staticmethod
def _model(modules, **attributes):
return SimpleNamespace(
named_modules=lambda: modules,
**attributes,
)
def test_supported_decoder_targets_are_accepted(self):
modules = [
(
"model.layers.0.qkv_proj",
ColumnParallelLinear.__new__(ColumnParallelLinear),
),
(
"model.layers.0.o_proj",
RowParallelLinear.__new__(RowParallelLinear),
),
(
"model.layers.0.fused_qkv_a_proj_with_mqa",
ReplicatedLinear.__new__(ReplicatedLinear),
),
]
self.backend.validate_lora_targets(
base_model=self._model(modules),
target_modules={
"qkv_proj",
"o_proj",
"fused_qkv_a_proj_with_mqa",
},
)
def test_unsupported_targets_are_rejected(self):
cases = {
"unknown decoder layer": (
self._model(
[
(
"model.layers.0.custom_proj",
torch.nn.Linear(2, 2),
)
]
),
{"custom_proj"},
"Linear",
),
"fused MoE": (
self._model(
[
(
"model.layers.0.mlp",
FusedMoE.__new__(FusedMoE),
)
]
),
{"gate_up_proj", "down_proj"},
"FusedMoE",
),
}
for name, (model, targets, expected) in cases.items():
with self.subTest(name=name), self.assertRaisesRegex(ValueError, expected):
self.backend.validate_lora_targets(
base_model=model,
target_modules=targets,
)
def test_nonoverlap_dense_calls_fall_back_to_triton(self):
x = object()
weights = object()
hidden = object()
base_output = object()
pruned_batch_info = object()
expected = object()
with (
patch.object(
TritonLoRABackend,
"run_lora_a_sgemm",
return_value=hidden,
) as run_lora_a,
patch.object(
TritonLoRABackend,
"run_lora_b_sgemm",
return_value=expected,
) as run_lora_b,
):
actual_hidden = self.backend.run_lora_a_sgemm(
x,
weights,
pruned_batch_info=pruned_batch_info,
)
actual = self.backend.run_lora_b_sgemm(
actual_hidden,
weights,
base_output=base_output,
pruned_batch_info=pruned_batch_info,
)
self.assertIs(actual_hidden, hidden)
self.assertIs(actual, expected)
run_lora_a.assert_called_once_with(
x,
weights,
pruned_batch_info,
1,
)
run_lora_b.assert_called_once_with(
hidden,
weights,
base_output,
pruned_batch_info,
)
def test_overlap_launch_selects_cublas(self):
pending = object()
x = object()
weights = object()
hidden = object()
base_output = object()
expected = object()
self.backend._pending_lora_a = pending
self.backend._consume_lora_a_overlap = MagicMock(return_value=hidden)
self.backend._run_lora_b = MagicMock(return_value=expected)
with (
patch.object(TritonLoRABackend, "run_lora_a_sgemm") as run_lora_a,
patch.object(TritonLoRABackend, "run_lora_b_sgemm") as run_lora_b,
):
actual_hidden = self.backend.run_lora_a_sgemm(x, weights)
actual = self.backend.run_lora_b_sgemm(
actual_hidden,
weights,
base_output=base_output,
)
self.assertIs(actual_hidden, hidden)
self.assertIs(actual, expected)
self.backend._consume_lora_a_overlap.assert_called_once_with(pending)
self.backend._run_lora_b.assert_called_once_with(
hidden,
weights,
base_output,
)
self.assertFalse(self.backend._use_cublas_lora_b)
run_lora_a.assert_not_called()
run_lora_b.assert_not_called()
def test_nonoverlap_qkv_call_falls_back_to_triton(self):
expected = object()
args = {
"x": object(),
"qkv_lora_a": object(),
"qkv_lora_b": object(),
"output_offset": object(),
"output_offset_cpu": object(),
"max_qkv_out_dim": 128,
"base_output": object(),
"n_slices": 2,
}
with patch.object(
TritonLoRABackend,
"run_qkv_lora",
return_value=expected,
) as run_qkv_lora:
actual = self.backend.run_qkv_lora(**args)
self.assertIs(actual, expected)
run_qkv_lora.assert_called_once_with(
args["x"],
args["qkv_lora_a"],
args["qkv_lora_b"],
args["output_offset"],
128,
args["base_output"],
2,
)
def test_manager_preflights_targets_before_wrapping(self):
manager = LoRAManager.__new__(LoRAManager)
manager.base_model = object()
manager.lora_backend = MagicMock()
manager._experts_shared_outer_override = None
manager.init_lora_adapters = MagicMock()
manager.init_lora_shapes = MagicMock(
side_effect=lambda **_: setattr(manager, "target_modules", {"qkv_proj"})
)
manager._detect_shared_outer_loras = MagicMock(return_value=False)
manager.init_lora_modules = MagicMock()
manager.init_memory_pool = MagicMock()
manager.update_lora_info = MagicMock()
manager.lora_backend.validate_lora_targets.side_effect = ValueError(
"unsupported target"
)
with self.assertRaisesRegex(ValueError, "unsupported target"):
manager.init_state(max_lora_rank=1, target_modules={"q_proj"})
manager.lora_backend.validate_lora_targets.assert_called_once_with(
base_model=manager.base_model,
target_modules={"qkv_proj"},
)
manager.init_lora_modules.assert_not_called()
if __name__ == "__main__":
unittest.main()
@@ -7,6 +7,7 @@ import torch
from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
from sglang.srt.runtime_context import get_context
from sglang.test.ci.ci_register import register_cpu_ci
@@ -178,17 +179,8 @@ class TestDecodeHiddenStateRetention(CustomTestCase):
second_step = torch.arange(16, dtype=torch.float32).view(8, 2)[4:]
def result(hidden_states):
return SimpleNamespace(
copy_done=None,
auxiliary_host_output=None,
routed_experts_output=None,
indexer_topk_output=None,
return GenerationBatchResult(
logits_output=SimpleNamespace(hidden_states=hidden_states),
next_token_ids=None,
can_run_cuda_graph=False,
num_correct_drafts=0,
num_block_accept_tokens=0,
num_cap_tokens=0,
speculative_num_draft_tokens=4,
)
@@ -10,6 +10,7 @@ from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.runtime_context import get_context
from sglang.srt.sampling.sampling_params import SamplingParams
from sglang.test.ci.ci_register import register_cpu_ci
@@ -78,17 +79,9 @@ def _make_processor() -> SchedulerBatchResultProcessor:
def _make_result():
return SimpleNamespace(
copy_done=None,
auxiliary_host_output=None,
routed_experts_output=None,
indexer_topk_output=None,
return GenerationBatchResult(
logits_output=SimpleNamespace(hidden_states=None, customized_info=None),
next_token_ids=[4],
can_run_cuda_graph=False,
num_correct_drafts=0,
num_block_accept_tokens=0,
num_cap_tokens=0,
speculative_num_draft_tokens=0,
)
@@ -14,6 +14,7 @@ from sglang.srt.managers.schedule_batch import Req
from sglang.srt.managers.scheduler_components.batch_result_processor import (
SchedulerBatchResultProcessor,
)
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.sampling.sampling_params import (
REQUEST_REASONING_END_TOKEN_IDS_KEY,
SamplingParams,
@@ -101,16 +102,10 @@ def _make_req(terminate_after: int) -> Req:
def _make_result(num_draft_tokens, accept_lens, flat_tokens):
return SimpleNamespace(
return GenerationBatchResult(
next_token_ids=torch.tensor(flat_tokens, dtype=torch.long),
accept_lens=torch.tensor(accept_lens, dtype=torch.long),
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,
copy_done=None,
grammar_advanced=False,
)
@@ -0,0 +1,67 @@
"""Scheduler containment for unsupported UNO requests."""
import unittest
from types import SimpleNamespace
from unittest.mock import MagicMock, patch
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel
maybe_stub_sgl_kernel()
from sglang.srt.disaggregation.utils import DisaggregationMode
from sglang.srt.managers.scheduler import Scheduler
register_cpu_ci(est_time=2, suite="base-a-test-cpu")
class TestSchedulerUnoRequestValidation(CustomTestCase):
def test_invalid_request_is_aborted_before_scheduler_admission(self):
scheduler = Scheduler.__new__(Scheduler)
scheduler.enable_session_radix_cache = False
scheduler.model_config = SimpleNamespace(
hf_eos_token_id={1},
vocab_size=128,
)
scheduler.disaggregation_mode = DisaggregationMode.NULL
scheduler.metrics_reporter = SimpleNamespace(enable_metrics=False)
scheduler.tokenizer = None
scheduler.dllm_config = None
scheduler._maybe_namespace_elastic_radix_cache = MagicMock()
scheduler.spec_algorithm = SimpleNamespace(
is_dflash_family=lambda: False,
is_uno=lambda: True,
)
scheduler.init_req_max_new_tokens = MagicMock()
scheduler._add_request_to_queue = MagicMock()
recv_req = MagicMock(
session_params=None,
session_id=None,
input_embeds=None,
bootstrap_port=1,
)
req = MagicMock()
error = "UNO request is unsupported."
with (
patch(
"sglang.srt.managers.scheduler.BeamCoordinator.request_beam_width",
return_value=1,
),
patch("sglang.srt.managers.scheduler.Req", return_value=req),
patch(
"sglang.srt.managers.scheduler.validate_uno_request",
return_value=error,
) as validate_uno_request,
):
scheduler.handle_generate_request(recv_req)
validate_uno_request.assert_called_once_with(req)
req.set_finish_with_abort.assert_called_once_with(error)
scheduler.init_req_max_new_tokens.assert_called_once_with(req)
scheduler._add_request_to_queue.assert_called_once_with(req)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,74 @@
"""CPU regressions for UNO aggregate token accounting."""
import unittest
from types import SimpleNamespace
from sglang.srt.managers.scheduler import Scheduler
from sglang.srt.managers.scheduler_components.metrics_reporter import (
SchedulerMetricsReporter,
)
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestUnoTokenAccounting(CustomTestCase):
def setUp(self):
self.batch_size = 2
self.result = GenerationBatchResult(
num_correct_drafts=3,
num_non_draft_tokens_per_req=2,
)
def test_generated_token_count_includes_both_non_draft_tokens(self):
self.assertEqual(self.result.get_num_generated_tokens(self.batch_size), 7)
self.assertEqual(
GenerationBatchResult(num_correct_drafts=3).get_num_generated_tokens(
self.batch_size
),
5,
)
def test_spec_metrics_keep_generated_and_draft_counts_separate(self):
reporter = SchedulerMetricsReporter.__new__(SchedulerMetricsReporter)
reporter.spec_num_accept_tokens = 0
reporter.spec_num_correct_drafts = 0
reporter.spec_num_forward_ct = 0
reporter.spec_num_block_accept_tokens = 0
reporter.spec_num_cap_tokens = 0
reporter.update_spec_metrics(
self.batch_size,
self.result.num_correct_drafts,
num_accept_tokens=self.result.get_num_generated_tokens(self.batch_size),
)
self.assertEqual(reporter.spec_num_accept_tokens, 7)
self.assertEqual(reporter.spec_num_correct_drafts, 3)
self.assertEqual(reporter.spec_num_forward_ct, 2)
def test_decode_moment_receives_full_generated_token_count(self):
scheduler = Scheduler.__new__(Scheduler)
scheduler._prev_step = (1, 10.0, False)
scheduler.decode_moment_totals = [0.0] * 6
batch = SimpleNamespace(
forward_mode=SimpleNamespace(
is_extend_without_speculative=lambda: False,
is_decode=lambda: True,
is_target_verify=lambda: False,
),
reqs=[SimpleNamespace(rid="req-0"), SimpleNamespace(rid="req-1")],
forward_iter=2,
launch_ts=10.001,
after_idle_gap=False,
)
scheduler._record_step_counters(batch, self.result)
self.assertEqual(scheduler.decode_moment_totals[5], 7)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,36 @@
"""Unit tests for UNO allocation sizing."""
import unittest
from sglang.srt.mem_cache.allocation_sizing import (
get_alloc_len_per_decode,
get_alloc_reserve_per_decode,
get_req_to_token_extra_context_len,
)
from sglang.srt.runtime_context import get_context, get_parallel
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestUnoAllocationSizing(CustomTestCase):
def test_page_size_one_row_covers_decode_reserve(self):
with (
get_context().override_server_args(
speculative_algorithm="UNO",
speculative_num_draft_tokens=8,
page_size=1,
),
get_parallel().override(attn_dcp_size=1),
):
self.assertEqual(get_alloc_len_per_decode(), 9)
self.assertEqual(get_alloc_reserve_per_decode(), 18)
self.assertGreaterEqual(
get_req_to_token_extra_context_len(),
get_alloc_reserve_per_decode(),
)
if __name__ == "__main__":
unittest.main()
@@ -14,6 +14,7 @@ import unittest
from unittest.mock import patch
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
@@ -25,6 +26,7 @@ class TestDraftRunnerSkipsLoRA(CustomTestCase):
runner = ModelRunner.__new__(ModelRunner)
runner.is_draft_worker = is_draft_worker
runner.lora_manager = None
runner.spec_algorithm = SpeculativeAlgorithm.NONE
with patch.object(ModelRunner, "init_lora_manager") as init_lora:
with patch("sglang.srt.model_executor.model_runner.get_lora") as get_lora:
get_lora.return_value.enable_lora = enable_lora
@@ -45,6 +45,10 @@ _DFLASH_DECODE = (
"speculative/dflash_info_v2.py",
"DFlashDraftInputV2.prepare_for_decode",
)
_UNO_DECODE = (
"speculative/uno_info.py",
"UnoDraftInput.prepare_for_decode",
)
_RESOLVE = (
"managers/scheduler_components/batch_result_processor.py",
"SchedulerBatchResultProcessor._resolve_spec_v2_tokens",
@@ -71,6 +75,8 @@ _OWNER_SITES = {
# one of these two owners for each speculative decode iteration.
(*_DFLASH_DECODE, "decode_batch_idx"): 1,
(*_DFLASH_DECODE, "evict"): 1,
(*_UNO_DECODE, "decode_batch_idx"): 1,
(*_UNO_DECODE, "evict"): 1,
(
"mem_cache/allocation.py",
"alloc_for_spec_decode",
@@ -0,0 +1,72 @@
"""CPU contracts for the fused suffix-attention merge dispatch guard."""
import unittest
from types import SimpleNamespace
import torch
from sglang.kernels.ops.attention.suffix_attention_merge import (
can_use_fused_suffix_attention_merge,
)
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestSuffixAttentionMergeDispatch(CustomTestCase):
def _inputs(self):
layer = SimpleNamespace(
head_dim=64,
v_head_dim=64,
is_cross_attention=False,
logit_cap=0.0,
)
q = torch.empty((16, 8 * 64), dtype=torch.bfloat16)
key_cache = torch.empty((8, 16, 2, 64), dtype=torch.bfloat16)
value_cache = torch.empty_like(key_cache)
return layer, q, key_cache, value_cache
def _eligible(self, **overrides):
layer, q, key_cache, value_cache = self._inputs()
arguments = dict(
layer=layer,
q=q,
key_cache=key_cache,
value_cache=value_cache,
extra_kwargs={},
)
arguments.update(overrides)
return can_use_fused_suffix_attention_merge(**arguments)
def test_standard_attention_is_eligible(self):
self.assertTrue(self._eligible())
def test_special_attention_features_fall_back(self):
self.assertFalse(self._eligible(extra_kwargs={"sinks": object()}))
layer, _, _, _ = self._inputs()
layer.is_cross_attention = True
self.assertFalse(self._eligible(layer=layer))
layer, _, _, _ = self._inputs()
layer.logit_cap = 20.0
self.assertFalse(self._eligible(layer=layer))
def test_unsupported_tensor_layout_falls_back(self):
layer, _, _, _ = self._inputs()
layer.v_head_dim = 32
self.assertFalse(self._eligible(layer=layer))
_, q, key_cache, value_cache = self._inputs()
self.assertFalse(
self._eligible(
q=q.float(),
key_cache=key_cache.float(),
value_cache=value_cache.float(),
)
)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,93 @@
"""Request-admission validation for UNO speculative decoding."""
import unittest
from types import SimpleNamespace
from sglang.srt.speculative.uno_validation import validate_uno_request
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
def _make_request(**overrides):
sampling_params = SimpleNamespace(
min_p=0.0,
json_schema=None,
regex=None,
ebnf=None,
structural_tag=None,
frequency_penalty=0.0,
presence_penalty=0.0,
repetition_penalty=1.0,
min_new_tokens=0,
logit_bias=None,
)
request = SimpleNamespace(
sampling_params=sampling_params,
grammar=None,
return_logprob=False,
return_hidden_states_mode=SimpleNamespace(need_capture=lambda: False),
custom_logit_processor=None,
lora_id=None,
)
for name, value in overrides.items():
target, field = name.split("__", maxsplit=1)
owner = sampling_params if target == "sampling_params" else request
setattr(owner, field, value)
return request
class TestUnoRequestValidation(CustomTestCase):
def test_supported_request_is_accepted(self):
self.assertIsNone(validate_uno_request(_make_request()))
def test_unsupported_request_features_are_rejected(self):
cases = {
"min_p": ({"sampling_params__min_p": 0.1}, "min_p"),
"grammar": ({"sampling_params__regex": "[0-9]+"}, "grammar"),
"logprobs": ({"request__return_logprob": True}, "logprobs"),
"hidden states": (
{
"request__return_hidden_states_mode": SimpleNamespace(
need_capture=lambda: True
)
},
"return_hidden_states",
),
"frequency penalty": (
{"sampling_params__frequency_penalty": 0.1},
"penalties",
),
"presence penalty": (
{"sampling_params__presence_penalty": 0.1},
"penalties",
),
"repetition penalty": (
{"sampling_params__repetition_penalty": 1.1},
"penalties",
),
"minimum new tokens": (
{"sampling_params__min_new_tokens": 1},
"penalties",
),
"logit bias": (
{"sampling_params__logit_bias": {1: 0.5}},
"logit_bias",
),
"custom processor": (
{"request__custom_logit_processor": "processor"},
"custom logit processors",
),
"public LoRA": ({"request__lora_id": "adapter"}, "LoRA"),
}
for name, (overrides, expected) in cases.items():
with self.subTest(name=name):
error = validate_uno_request(_make_request(**overrides))
self.assertIsNotNone(error)
self.assertIn(expected, error)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,73 @@
"""Startup validation for UNO configuration."""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
from sglang.srt.arg_groups.speculative_hook import _handle_uno
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestUnoTreeConfig(CustomTestCase):
def test_unsupported_runtime_modes_are_rejected_at_startup(self):
cases = {
"deterministic inference": (
{"enable_deterministic_inference": True},
"enable-deterministic-inference",
),
"strict thinking": (
{"enable_strict_thinking": True},
"enable-strict-thinking",
),
}
for name, (overrides, expected) in cases.items():
with self.subTest(name=name):
values = {
"device": "cuda",
"speculative_draft_model_path": None,
"uno_lora_path": "/tmp/uno-lora",
"enable_deterministic_inference": False,
"enable_strict_thinking": False,
}
values.update(overrides)
server_args = SimpleNamespace(**values)
with (
patch(
"sglang.srt.arg_groups.speculative_hook.resolving_view",
side_effect=lambda args: args,
),
self.assertRaisesRegex(ValueError, expected),
):
_handle_uno(server_args)
def test_parent_list_overflow_is_rejected_at_startup(self):
"""An invalid tree must not survive startup and crash on first decode."""
server_args = SimpleNamespace(
device="cuda",
enable_deterministic_inference=False,
enable_strict_thinking=False,
speculative_draft_model_path=None,
uno_lora_path="/tmp/uno-lora",
speculative_num_draft_tokens=8,
speculative_num_steps=3,
speculative_eagle_topk=2,
)
with (
patch(
"sglang.srt.arg_groups.speculative_hook.resolving_view",
side_effect=lambda args: args,
),
patch("sglang.srt.arg_groups.speculative_hook.declare_resolution"),
self.assertRaisesRegex(ValueError, "parent-list ABI"),
):
_handle_uno(server_args)
if __name__ == "__main__":
unittest.main()
@@ -0,0 +1,105 @@
"""CPU contracts for UNO tree compact target sampling."""
import unittest
from types import SimpleNamespace
from unittest.mock import patch
import torch
from sglang.srt.speculative.eagle_utils import (
_can_use_sparse_uno_tree_target_sampling,
)
from sglang.srt.speculative.uno_utils import sample_uno_tree_target_tokens
from sglang.test.ci.ci_register import register_cpu_ci
from sglang.test.test_utils import CustomTestCase
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
class TestUnoTreeSparseSampling(CustomTestCase):
def test_sparse_dispatch_guard(self):
sampling_info = SimpleNamespace(
sampling_seed=None,
need_min_p_sampling=False,
)
spec_config = SimpleNamespace(
speculative_use_rejection_sampling=False,
)
with (
patch("sglang.srt.speculative.eagle_utils._is_cuda", True),
patch(
"sglang.srt.speculative.eagle_utils.get_spec",
return_value=spec_config,
),
):
self.assertTrue(
_can_use_sparse_uno_tree_target_sampling(128, sampling_info)
)
self.assertFalse(
_can_use_sparse_uno_tree_target_sampling(None, sampling_info)
)
self.assertFalse(
_can_use_sparse_uno_tree_target_sampling(129, sampling_info)
)
sampling_info.sampling_seed = torch.tensor([1])
self.assertFalse(
_can_use_sparse_uno_tree_target_sampling(128, sampling_info)
)
sampling_info.sampling_seed = None
sampling_info.need_min_p_sampling = True
self.assertFalse(
_can_use_sparse_uno_tree_target_sampling(128, sampling_info)
)
sampling_info.need_min_p_sampling = False
spec_config.speculative_use_rejection_sampling = True
self.assertFalse(
_can_use_sparse_uno_tree_target_sampling(128, sampling_info)
)
def test_targets_are_sampled_from_compact_support(self):
support_ids = torch.tensor(
[
[[10, 11], [20, 21], [30, 31]],
[[40, 41], [50, 51], [60, 61]],
],
dtype=torch.int64,
)
support_probs = torch.full((2, 3, 2), 0.5)
sampled_offsets = torch.tensor(
[[0], [1], [0], [1], [0], [1]],
dtype=torch.long,
)
sampling_info = SimpleNamespace()
next_token_logits = torch.empty((6, 100))
with (
patch(
"sglang.srt.speculative.uno_utils._build_sparse_target_support",
return_value=(support_ids, support_probs),
) as build_support,
patch(
"sglang.srt.speculative.uno_utils.fast_sample",
return_value=(torch.empty((6, 1)), sampled_offsets),
),
):
targets = sample_uno_tree_target_tokens(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
batch_size=2,
verify_width=3,
max_top_k=2,
)
self.assertEqual(targets.tolist(), [[10, 21, 30], [41, 50, 61]])
build_support.assert_called_once_with(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
batch_size=2,
forward_width=3,
max_top_k=2,
)
if __name__ == "__main__":
unittest.main()