Support EAGLE speculative decoding in the KV-canary (#26813)

This commit is contained in:
fzyzcjy
2026-05-31 09:57:01 +08:00
committed by GitHub
parent 30a22cc360
commit 9c43c3719f
7 changed files with 131 additions and 11 deletions
+8 -2
View File
@@ -35,11 +35,14 @@ def install_canary(
) )
device = torch.device(model_runner.device) device = torch.device(model_runner.device)
# EAGLE draft worker pools rotate input_ids so slot ``p`` stores K/V for the token at position ``p+1``;
# target pools have no such shift. Threaded into the plan-side expected-token gather kernel.
kv_token_id_vs_position_offset = 1 if model_runner.is_draft_worker else 0
buffer_groups = attach_canary_buffers( buffer_groups = attach_canary_buffers(
pool=model_runner.token_to_kv_pool, pool=model_runner.token_to_kv_pool,
config=config, config=config,
device=device, device=device,
kv_token_id_vs_position_offset=0, kv_token_id_vs_position_offset=kv_token_id_vs_position_offset,
) )
launch_capacities = CanaryLaunchCapacities.from_args( launch_capacities = CanaryLaunchCapacities.from_args(
server_args=model_runner.server_args, server_args=model_runner.server_args,
@@ -48,6 +51,7 @@ def install_canary(
pool_slot_count=model_runner.max_total_num_tokens, pool_slot_count=model_runner.max_total_num_tokens,
) )
swa_window_size = model_runner.sliding_window_size or 0 swa_window_size = model_runner.sliding_window_size or 0
speculative_num_steps = int(server_args.speculative_num_steps or 1)
manager = CanaryManager( manager = CanaryManager(
config=config, config=config,
buffer_groups=buffer_groups, buffer_groups=buffer_groups,
@@ -55,6 +59,7 @@ def install_canary(
req_to_token_pool=model_runner.req_to_token_pool, req_to_token_pool=model_runner.req_to_token_pool,
launch_capacities=launch_capacities, launch_capacities=launch_capacities,
swa_window_size=swa_window_size, swa_window_size=swa_window_size,
speculative_num_steps=speculative_num_steps,
) )
_patch_model_forward(model_runner=model_runner, manager=manager) _patch_model_forward(model_runner=model_runner, manager=manager)
@@ -64,13 +69,14 @@ def install_canary(
logger.info( logger.info(
"install_canary: disaggregation_mode=%s config=%s " "install_canary: disaggregation_mode=%s config=%s "
"launch_capacities=%s n_buffer_groups=%d buffer_group_kinds=%s " "launch_capacities=%s n_buffer_groups=%d buffer_group_kinds=%s "
"swa_window_size=%d", "swa_window_size=%d speculative_num_steps=%d",
server_args.disaggregation_mode, server_args.disaggregation_mode,
config, config,
launch_capacities, launch_capacities,
len(buffer_groups), len(buffer_groups),
[g.kind.name for g in buffer_groups], [g.kind.name for g in buffer_groups],
swa_window_size, swa_window_size,
speculative_num_steps,
) )
return manager return manager
@@ -42,6 +42,7 @@ class CanaryManager:
req_to_token_pool: "ReqToTokenPool", req_to_token_pool: "ReqToTokenPool",
launch_capacities: CanaryLaunchCapacities, launch_capacities: CanaryLaunchCapacities,
swa_window_size: int = 0, swa_window_size: int = 0,
speculative_num_steps: int = 1,
) -> None: ) -> None:
self.config = config self.config = config
self._req_to_token_pool = req_to_token_pool self._req_to_token_pool = req_to_token_pool
@@ -92,7 +93,8 @@ class CanaryManager:
swa_window_size=self._swa_window_size, swa_window_size=self._swa_window_size,
outer_step_counter_getter=self._get_outer_step_counter, outer_step_counter_getter=self._get_outer_step_counter,
) )
self._single_forward_managers: tuple[SingleForwardManager, ...] = ( num_sfms = max(1, speculative_num_steps - 1)
self._single_forward_managers: tuple[SingleForwardManager, ...] = tuple(
SingleForwardManager( SingleForwardManager(
config=config, config=config,
device=device, device=device,
@@ -105,7 +107,8 @@ class CanaryManager:
per_forward_write_req_capacity=launch_capacities.per_forward_write_req_capacity, per_forward_write_req_capacity=launch_capacities.per_forward_write_req_capacity,
per_forward_write_entry_capacity=launch_capacities.per_forward_write_entry_capacity, per_forward_write_entry_capacity=launch_capacities.per_forward_write_entry_capacity,
d2h_stream=self._d2h_stream, d2h_stream=self._d2h_stream,
), )
for _ in range(num_sfms)
) )
@contextlib.contextmanager @contextlib.contextmanager
@@ -434,7 +434,11 @@ class EAGLEDraftExtendCudaGraphRunner:
) )
self.deepep_adapter.capture(is_extend_in_batch=True) self.deepep_adapter.capture(is_extend_in_batch=True)
canary_ctx = contextlib.nullcontext() canary_ctx = (
c.with_active_single_forward_manager(0)
if (c := self.model_runner.canary_manager) is not None
else contextlib.nullcontext()
)
with canary_ctx: with canary_ctx:
self._capture_init(run_once) self._capture_init(run_once)
@@ -12,6 +12,7 @@ from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_extend_npu_graph_r
from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner import ( from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_npu_graph_runner import (
EAGLEDraftNpuGraphRunner, EAGLEDraftNpuGraphRunner,
) )
from sglang.srt.kv_canary.runner.canary_manager import context_tuple
from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend from sglang.srt.layers.attention.tokenspeed_mla_backend import TokenspeedMLABackend
from sglang.srt.layers.attention.triton_backend import TritonAttnBackend from sglang.srt.layers.attention.triton_backend import TritonAttnBackend
from sglang.srt.layers.attention.trtllm_mla_backend import ( from sglang.srt.layers.attention.trtllm_mla_backend import (
@@ -205,6 +206,9 @@ class EagleDraftWorker(BaseDraftWorker):
) )
self.init_cuda_graphs() self.init_cuda_graphs()
if (c := self.draft_runner.canary_manager) is not None:
c.mark_init_finished()
self.tree_mask_mode = TreeMaskMode.FULL_MASK self.tree_mask_mode = TreeMaskMode.FULL_MASK
self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device) self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device)
@@ -363,7 +367,15 @@ class EagleDraftWorker(BaseDraftWorker):
self.speculative_num_steps, self.speculative_num_steps,
) )
canary_outside_ctx = contextlib.nullcontext() n_inner = self.speculative_num_steps - 1
canary_outside_ctx = (
c.with_ops_outside_graph(
single_forward_indices=list(range(n_inner)),
maybe_inaccurate_forward_batch=forward_batch,
)
if (c := self.draft_runner.canary_manager) is not None
else contextlib.nullcontext()
)
with canary_outside_ctx: with canary_outside_ctx:
# Run draft # Run draft
@@ -501,9 +513,13 @@ class EagleDraftWorker(BaseDraftWorker):
spec_info.hidden_states = hidden_states spec_info.hidden_states = hidden_states
# Run forward under a per-step ForwardContext so the model layer # Run forward under a per-step ForwardContext so the model layer
# reads attn_backends[i] for the i-th draft step. ``_forward_raw`` # reads attn_backends[i] for the i-th draft step, plus a canary
# honors the outer context and does not override. # index context so canary tracks which draft step is active.
canary_index_ctx = contextlib.nullcontext() canary_index_ctx = (
c.with_active_single_forward_manager(i)
if (c := self.draft_runner.canary_manager) is not None
else contextlib.nullcontext()
)
with forward_context( with forward_context(
ForwardContext(attn_backend=self.draft_attn_backend.attn_backends[i]) ForwardContext(attn_backend=self.draft_attn_backend.attn_backends[i])
), canary_index_ctx: ), canary_index_ctx:
@@ -619,7 +635,17 @@ class EagleDraftWorker(BaseDraftWorker):
if mm_input_embeds is not None: if mm_input_embeds is not None:
forward_batch.mm_input_embeds = mm_input_embeds forward_batch.mm_input_embeds = mm_input_embeds
canary_ctx = contextlib.nullcontext() canary_ctx = (
context_tuple(
c.with_ops_outside_graph(
single_forward_indices=[0],
maybe_inaccurate_forward_batch=forward_batch,
),
c.with_active_single_forward_manager(0),
)
if (c := self.draft_runner.canary_manager) is not None
else contextlib.nullcontext()
)
with canary_ctx: with canary_ctx:
logits_output = self.draft_runner.forward(forward_batch).logits_output logits_output = self.draft_runner.forward(forward_batch).logits_output
maybe_detect_nan(logits_output.next_token_logits, "draft_extend_for_prefill") maybe_detect_nan(logits_output.next_token_logits, "draft_extend_for_prefill")
@@ -676,7 +702,17 @@ class EagleDraftWorker(BaseDraftWorker):
and self.cuda_graph_runner_for_draft_extend.can_run(forward_batch) and self.cuda_graph_runner_for_draft_extend.can_run(forward_batch)
) )
canary_ctx = contextlib.nullcontext() canary_ctx = (
context_tuple(
c.with_ops_outside_graph(
single_forward_indices=[0],
maybe_inaccurate_forward_batch=forward_batch,
),
c.with_active_single_forward_manager(0),
)
if (c := self.draft_runner.canary_manager) is not None
else contextlib.nullcontext()
)
with canary_ctx: with canary_ctx:
if can_cuda_graph: if can_cuda_graph:
draft_logits_output = self.cuda_graph_runner_for_draft_extend.replay( draft_logits_output = self.cuda_graph_runner_for_draft_extend.replay(
@@ -49,6 +49,7 @@ def make_manager(
group: CanaryBufferGroup | None = None, group: CanaryBufferGroup | None = None,
req_pool: SimpleNamespace | None = None, req_pool: SimpleNamespace | None = None,
per_forward_verify_capacity: int = 16, per_forward_verify_capacity: int = 16,
speculative_num_steps: int = 1,
) -> CanaryManager: ) -> CanaryManager:
if config is None: if config is None:
config = make_config() config = make_config()
@@ -66,6 +67,7 @@ def make_manager(
per_forward_write_req_capacity=2, per_forward_write_req_capacity=2,
per_forward_write_entry_capacity=8, per_forward_write_entry_capacity=8,
), ),
speculative_num_steps=speculative_num_steps,
) )
@@ -255,5 +255,40 @@ def _drive_one_cycle(manager, forward_batch) -> None:
manager.post_ops_maybe_inside_graph(forward_batch, pre_ops_output) manager.post_ops_maybe_inside_graph(forward_batch, pre_ops_output)
class TestCanaryManagerActiveSingleForwardManagerDispatch(CanaryManagerTestCase):
def test_pre_ops_maybe_inside_graph_dispatches_to_bracketed_sfm(
self,
) -> None:
"""Verify the dispatcher routes phase 2 to the bracketed SingleForwardManager."""
manager = make_manager(device=self.device, speculative_num_steps=3)
forward_batch = make_forward_batch(self.device)
target_sfm = manager._single_forward_managers[1]
observed: list[object] = []
original_phase_2 = target_sfm.pre_ops_maybe_inside_graph
def _record(fb):
observed.append(fb)
return original_phase_2(fb)
target_sfm.pre_ops_maybe_inside_graph = _record
manager._single_forward_managers[0].pre_ops_outside_graph(
maybe_inaccurate_forward_batch=forward_batch
)
manager._single_forward_managers[1].pre_ops_outside_graph(
maybe_inaccurate_forward_batch=forward_batch
)
with manager.with_active_single_forward_manager(1):
manager.pre_ops_maybe_inside_graph(forward_batch)
self.assertEqual(observed, [forward_batch])
def test_pre_ops_maybe_inside_graph_asserts_outside_bracket(
self,
) -> None:
manager = make_manager(device=self.device)
forward_batch = make_forward_batch(self.device)
with self.assertRaises(AssertionError):
manager.pre_ops_maybe_inside_graph(forward_batch)
if __name__ == "__main__": if __name__ == "__main__":
unittest.main() unittest.main()
@@ -0,0 +1,34 @@
from __future__ import annotations
import unittest
from sglang.test.ci.ci_register import register_cuda_ci
from sglang.test.mock_model.utils import MOCK_MODEL_PATH, run_mock_model_bench_serving
from sglang.test.test_utils import CustomTestCase
register_cuda_ci(est_time=600, stage="extra-a", runner_config="1-gpu-small")
class TestE2ESpeculativeEagle(CustomTestCase):
def test_spec_eagle_no_canary_violation(self) -> None:
run_mock_model_bench_serving(
extra_server_args=[
"--speculative-algorithm",
"EAGLE",
"--speculative-draft-model-path",
MOCK_MODEL_PATH,
"--speculative-num-steps",
"1",
"--speculative-eagle-topk",
"1",
"--speculative-num-draft-tokens",
"2",
"--mem-fraction-static",
"0.45",
],
input_check_enabled=False,
)
if __name__ == "__main__":
unittest.main()