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)
# 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(
pool=model_runner.token_to_kv_pool,
config=config,
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(
server_args=model_runner.server_args,
@@ -48,6 +51,7 @@ def install_canary(
pool_slot_count=model_runner.max_total_num_tokens,
)
swa_window_size = model_runner.sliding_window_size or 0
speculative_num_steps = int(server_args.speculative_num_steps or 1)
manager = CanaryManager(
config=config,
buffer_groups=buffer_groups,
@@ -55,6 +59,7 @@ def install_canary(
req_to_token_pool=model_runner.req_to_token_pool,
launch_capacities=launch_capacities,
swa_window_size=swa_window_size,
speculative_num_steps=speculative_num_steps,
)
_patch_model_forward(model_runner=model_runner, manager=manager)
@@ -64,13 +69,14 @@ def install_canary(
logger.info(
"install_canary: disaggregation_mode=%s config=%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,
config,
launch_capacities,
len(buffer_groups),
[g.kind.name for g in buffer_groups],
swa_window_size,
speculative_num_steps,
)
return manager
@@ -42,6 +42,7 @@ class CanaryManager:
req_to_token_pool: "ReqToTokenPool",
launch_capacities: CanaryLaunchCapacities,
swa_window_size: int = 0,
speculative_num_steps: int = 1,
) -> None:
self.config = config
self._req_to_token_pool = req_to_token_pool
@@ -92,7 +93,8 @@ class CanaryManager:
swa_window_size=self._swa_window_size,
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(
config=config,
device=device,
@@ -105,7 +107,8 @@ class CanaryManager:
per_forward_write_req_capacity=launch_capacities.per_forward_write_req_capacity,
per_forward_write_entry_capacity=launch_capacities.per_forward_write_entry_capacity,
d2h_stream=self._d2h_stream,
),
)
for _ in range(num_sfms)
)
@contextlib.contextmanager
@@ -434,7 +434,11 @@ class EAGLEDraftExtendCudaGraphRunner:
)
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:
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 (
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.triton_backend import TritonAttnBackend
from sglang.srt.layers.attention.trtllm_mla_backend import (
@@ -205,6 +206,9 @@ class EagleDraftWorker(BaseDraftWorker):
)
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.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device)
@@ -363,7 +367,15 @@ class EagleDraftWorker(BaseDraftWorker):
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:
# Run draft
@@ -501,9 +513,13 @@ class EagleDraftWorker(BaseDraftWorker):
spec_info.hidden_states = hidden_states
# Run forward under a per-step ForwardContext so the model layer
# reads attn_backends[i] for the i-th draft step. ``_forward_raw``
# honors the outer context and does not override.
canary_index_ctx = contextlib.nullcontext()
# reads attn_backends[i] for the i-th draft step, plus a canary
# index context so canary tracks which draft step is active.
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(
ForwardContext(attn_backend=self.draft_attn_backend.attn_backends[i])
), canary_index_ctx:
@@ -619,7 +635,17 @@ class EagleDraftWorker(BaseDraftWorker):
if mm_input_embeds is not None:
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:
logits_output = self.draft_runner.forward(forward_batch).logits_output
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)
)
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:
if can_cuda_graph:
draft_logits_output = self.cuda_graph_runner_for_draft_extend.replay(
@@ -49,6 +49,7 @@ def make_manager(
group: CanaryBufferGroup | None = None,
req_pool: SimpleNamespace | None = None,
per_forward_verify_capacity: int = 16,
speculative_num_steps: int = 1,
) -> CanaryManager:
if config is None:
config = make_config()
@@ -66,6 +67,7 @@ def make_manager(
per_forward_write_req_capacity=2,
per_forward_write_entry_capacity=8,
),
speculative_num_steps=speculative_num_steps,
)