diff --git a/python/sglang/srt/speculative/adaptive_runtime_state.py b/python/sglang/srt/speculative/adaptive_runtime_state.py index fa9211dd3..d28e32406 100644 --- a/python/sglang/srt/speculative/adaptive_runtime_state.py +++ b/python/sglang/srt/speculative/adaptive_runtime_state.py @@ -1,18 +1,10 @@ from dataclasses import dataclass from typing import TYPE_CHECKING, Protocol -from sglang.srt.speculative.adaptive_spec_params import AdaptiveSpeculativeParams - if TYPE_CHECKING: from sglang.srt.layers.attention.base_attn_backend import AttentionBackend from sglang.srt.model_executor.cpu_graph_runner import CPUGraphRunner from sglang.srt.model_executor.runner import DecodeCudaGraphRunner - from sglang.srt.speculative.eagle_draft_cuda_graph_runner import ( - EAGLEDraftCudaGraphRunner, - ) - from sglang.srt.speculative.eagle_draft_extend_cuda_graph_runner import ( - EAGLEDraftExtendCudaGraphRunner, - ) @dataclass @@ -20,10 +12,9 @@ class SpecRuntimeState: """A complete set of runtime resources bound to a specific speculative decoding configuration. - Each decode round runs three stages — draft, verify, extend — and every - stage has shape-dependent resources (attention backends and CUDA graphs) - that must match the current configuration. Switching adaptive steps - means swapping the entire state atomically. + The draft and verify resources are required by adaptive workers. Algorithms + with a draft-extend stage can also populate its optional resources. + Switching adaptive steps swaps the entire state atomically. """ # -- Configuration (determines shapes for all stages) -- @@ -32,7 +23,7 @@ class SpecRuntimeState: # -- Draft stage: draft model multi-step autoregressive generation -- draft_attn_backend: "AttentionBackend | None" - cuda_graph_runner: "EAGLEDraftCudaGraphRunner | None" + cuda_graph_runner: "DecodeCudaGraphRunner | None" # -- Verify stage: target model one-pass tree verification -- target_attn_backend: "AttentionBackend" @@ -40,7 +31,7 @@ class SpecRuntimeState: # -- Extend stage: draft model KV cache catch-up after verify -- draft_extend_attn_backend: "AttentionBackend | None" - cuda_graph_runner_for_draft_extend: "EAGLEDraftExtendCudaGraphRunner | None" + cuda_graph_runner_for_draft_extend: "DecodeCudaGraphRunner | None" class AdaptiveSpecWorker(Protocol): @@ -58,6 +49,23 @@ class AdaptiveSpecWorker(Protocol): def apply_runtime_state(self, state: SpecRuntimeState) -> None: ... +class AdaptiveSpecPolicy(Protocol): + """Policy interface used by AdaptiveController to select runtime states.""" + + @property + def candidate_steps(self) -> list[int]: ... + + def set_cuda_graph_bs(self, cuda_graph_bs: list[int] | None) -> None: ... + + def get_steps_for_batch(self, batch_size: int) -> int: ... + + def on_verify_complete( + self, num_correct_drafts_per_req: list[int], batch_size: int + ) -> int | None: ... + + def cuda_graph_bs_for_step(self, step: int) -> list[int] | None: ... + + class AdaptiveController: """Facade that owns adaptive decision-making and runtime state switching. @@ -71,12 +79,13 @@ class AdaptiveController: 2. Call on_verify_complete(num_correct_drafts_per_req) after each decode verify. """ - def __init__(self, worker: AdaptiveSpecWorker, config_path: str | None = None): + def __init__( + self, + worker: AdaptiveSpecWorker, + policy: AdaptiveSpecPolicy, + ): self.worker = worker - self.params = AdaptiveSpeculativeParams( - initial_steps=worker.speculative_num_steps, - cfg_path=config_path, - ) + self.params: AdaptiveSpecPolicy = policy self._states: dict[int, SpecRuntimeState] = {} @property @@ -118,7 +127,7 @@ class AdaptiveController: def on_verify_complete( self, num_correct_drafts_per_req: list[int], batch_size: int ) -> None: - """Feed verify results; switch runtime state if EMA warrants it.""" + """Feed verify results; switch runtime state if the policy requests it.""" new_step = self.params.on_verify_complete( num_correct_drafts_per_req, batch_size ) diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index afdd57ac7..875c81688 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -60,6 +60,7 @@ from sglang.srt.speculative.adaptive_runtime_state import ( AdaptiveController, SpecRuntimeState, ) +from sglang.srt.speculative.adaptive_spec_params import AdaptiveSpeculativeParams from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase from sglang.srt.speculative.draft_utils import DraftBackendFactory from sglang.srt.speculative.eagle_draft_cuda_graph_runner import ( @@ -1101,7 +1102,10 @@ class EAGLEWorkerV2(BaseSpecWorker): if get_spec().speculative_adaptive and self._hosts_draft: self.adaptive_controller = AdaptiveController( self, - config_path=get_spec().speculative_adaptive_config, + AdaptiveSpeculativeParams( + initial_steps=self.speculative_num_steps, + cfg_path=get_spec().speculative_adaptive_config, + ), ) # Some dummy tensors diff --git a/test/registered/unit/spec/test_adaptive_runtime_state.py b/test/registered/unit/spec/test_adaptive_runtime_state.py new file mode 100644 index 000000000..5d7891395 --- /dev/null +++ b/test/registered/unit/spec/test_adaptive_runtime_state.py @@ -0,0 +1,132 @@ +import unittest + +from sglang.srt.speculative.adaptive_runtime_state import ( + AdaptiveController, + SpecRuntimeState, +) +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +def _make_state(steps: int) -> SpecRuntimeState: + return SpecRuntimeState( + speculative_num_steps=steps, + speculative_num_draft_tokens=steps + 1, + draft_attn_backend=None, + cuda_graph_runner=None, + target_attn_backend=None, + target_graph_runner=None, + draft_extend_attn_backend=None, + cuda_graph_runner_for_draft_extend=None, + ) + + +class _FakeWorker: + def __init__(self, initial_steps: int = 3): + self.speculative_num_steps = initial_steps + self.build_calls = [] + self.applied_steps = [] + + def build_adaptive_runtime_state( + self, + speculative_num_steps: int, + speculative_num_draft_tokens: int, + cuda_graph_bs: list[int] | None = None, + ) -> SpecRuntimeState: + self.build_calls.append( + (speculative_num_steps, speculative_num_draft_tokens, cuda_graph_bs) + ) + return _make_state(speculative_num_steps) + + def apply_runtime_state(self, state: SpecRuntimeState) -> None: + self.speculative_num_steps = state.speculative_num_steps + self.applied_steps.append(state.speculative_num_steps) + + +class _FakePolicy: + def __init__(self): + self.candidate_steps = [1, 3] + self.cuda_graph_bs = None + self.feedback_step = None + + def set_cuda_graph_bs(self, cuda_graph_bs: list[int] | None) -> None: + self.cuda_graph_bs = cuda_graph_bs + + def get_steps_for_batch(self, batch_size: int) -> int: + return 1 if batch_size >= 8 else 3 + + def on_verify_complete( + self, num_correct_drafts_per_req: list[int], batch_size: int + ) -> int | None: + return self.feedback_step + + def cuda_graph_bs_for_step(self, step: int) -> list[int] | None: + if self.cuda_graph_bs is None: + return None + return [batch_size for batch_size in self.cuda_graph_bs if batch_size <= step] + + +class TestAdaptiveController(unittest.TestCase): + def test_injected_policy_builds_pruned_states_and_applies_initial_state(self): + worker = _FakeWorker(initial_steps=3) + policy = _FakePolicy() + controller = AdaptiveController(worker, policy) + + controller.init_states(cuda_graph_bs=[1, 2, 4]) + + self.assertIs(controller.params, policy) + self.assertEqual(policy.cuda_graph_bs, [1, 2, 4]) + self.assertEqual( + worker.build_calls, + [ + (1, 2, [1]), + (3, 4, [1, 2]), + ], + ) + self.assertEqual(worker.applied_steps, [3]) + + def test_registered_initial_state_is_reused(self): + worker = _FakeWorker(initial_steps=3) + controller = AdaptiveController(worker, _FakePolicy()) + controller.register(_make_state(3)) + + controller.init_states() + + self.assertEqual(worker.build_calls, [(1, 2, None)]) + self.assertEqual(worker.applied_steps, [3]) + + def test_batch_activation_is_idempotent(self): + worker = _FakeWorker(initial_steps=3) + controller = AdaptiveController(worker, _FakePolicy()) + controller.init_states() + + controller.activate_step_by_batch(batch_size=8) + controller.activate_step_by_batch(batch_size=8) + self.assertEqual(worker.applied_steps, [3, 1]) + + def test_verify_feedback_can_activate_a_state(self): + worker = _FakeWorker(initial_steps=3) + policy = _FakePolicy() + controller = AdaptiveController(worker, policy) + controller.init_states() + + controller.on_verify_complete([1], batch_size=1) + policy.feedback_step = 1 + controller.on_verify_complete([1], batch_size=1) + self.assertEqual(worker.applied_steps, [3, 1]) + + def test_missing_state_fails_loud(self): + worker = _FakeWorker(initial_steps=3) + controller = AdaptiveController(worker, _FakePolicy()) + controller.init_states() + del controller._states[1] + + with self.assertRaisesRegex( + ValueError, "Missing adaptive runtime state for steps=1" + ): + controller.activate_step_by_batch(batch_size=8) + + +if __name__ == "__main__": + unittest.main()