Support EAGLE speculative decoding in the KV-canary (#26813)
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user