[Spec] Fix EagleDraftWorker draft-extend attn backend assignment (#28096)

Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
Cheng Wan
2026-06-12 15:47:05 -07:00
committed by GitHub
co-authored by Claude Opus 4.8
parent fd977adbd6
commit d601edab73
2 changed files with 82 additions and 0 deletions
@@ -322,6 +322,8 @@ class EagleDraftWorker(BaseDraftWorker):
)
self.draft_runner.draft_attn_backend = self.draft_attn_backend
if self.draft_extend_attn_backend is not None:
self.draft_runner.attn_backend = self.draft_extend_attn_backend
self.tree_mask_mode = TreeMaskMode.FULL_MASK
def init_cuda_graphs(self):
@@ -1115,6 +1117,12 @@ class EAGLEWorkerV2(BaseSpecWorker):
dw.draft_runner.draft_attn_backend = state.draft_attn_backend
dw.cuda_graph_runner = state.cuda_graph_runner
dw.draft_extend_attn_backend = state.draft_extend_attn_backend
# Keep the runner's attn_backend in step with the active draft-extend
# backend (the draft-extend forward reads draft_runner.attn_backend);
# mirrors init_attention_backend. When None, the runner keeps its
# initialized backend (consistent across step configs).
if state.draft_extend_attn_backend is not None:
dw.draft_runner.attn_backend = state.draft_extend_attn_backend
dw.cuda_graph_runner_for_draft_extend = state.cuda_graph_runner_for_draft_extend
dw._rebuild_topk1_chain_buffers()
@@ -1148,6 +1156,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
dw.draft_attn_backend,
dw.draft_extend_attn_backend,
dw.draft_runner.draft_attn_backend,
dw.draft_runner.attn_backend,
dw.cuda_graph_runner,
dw.cuda_graph_runner_for_draft_extend,
sa.speculative_num_steps,
@@ -1183,6 +1192,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
dw.draft_attn_backend,
dw.draft_extend_attn_backend,
dw.draft_runner.draft_attn_backend,
dw.draft_runner.attn_backend,
dw.cuda_graph_runner,
dw.cuda_graph_runner_for_draft_extend,
sa.speculative_num_steps,
@@ -12,6 +12,7 @@ from unittest.mock import patch
import torch
from sglang.srt.speculative.adaptive_runtime_state import SpecRuntimeState
from sglang.srt.speculative.eagle_utils import organize_draft_results
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
from sglang.srt.utils import get_device
@@ -150,6 +151,77 @@ class TestEagleWorkerV2BackendFallback(CustomTestCase):
self.assertIs(worker.draft_runner.draft_attn_backend, decode_backend)
self.assertIs(worker.draft_runner.attn_backend, draft_extend_backend)
def _make_adaptive_worker(self, runner_attn_backend):
"""An EAGLEWorkerV2 with a draft worker whose state-machine fields are
filled with sentinels, sufficient to drive _override_worker_state /
apply_runtime_state without touching the GPU."""
draft_runner = SimpleNamespace(
draft_attn_backend=object(),
attn_backend=runner_attn_backend,
)
draft_worker = SimpleNamespace(
speculative_num_steps=2,
speculative_num_draft_tokens=3,
draft_attn_backend=object(),
draft_extend_attn_backend=object(),
cuda_graph_runner=object(),
cuda_graph_runner_for_draft_extend=object(),
draft_runner=draft_runner,
# _override_worker_state / apply_runtime_state call this hook; the
# topk=1 buffers are exercised by the fast-path tests above.
_rebuild_topk1_chain_buffers=lambda: None,
)
worker = object.__new__(EAGLEWorkerV2)
worker._draft_worker = draft_worker
worker._target_worker = SimpleNamespace(
model_runner=SimpleNamespace(
attn_backend=object(), decode_cuda_graph_runner=object()
)
)
worker.speculative_num_steps = 2
worker.speculative_num_draft_tokens = 3
worker.server_args = SimpleNamespace(
speculative_num_steps=2,
speculative_num_draft_tokens=3,
cuda_graph_bs_decode=None,
disable_cuda_graph=False,
)
return worker, draft_worker
def test_override_worker_state_restores_runner_attn_backend(self):
# build_adaptive_runtime_state runs init_attention_backend inside this
# context for each candidate step; the runner backend it assigns must
# not leak into the live worker.
initial_backend = object()
candidate_backend = object()
worker, dw = self._make_adaptive_worker(initial_backend)
with worker._override_worker_state(3, 4):
dw.draft_runner.attn_backend = candidate_backend
self.assertIs(dw.draft_runner.attn_backend, candidate_backend)
self.assertIs(dw.draft_runner.attn_backend, initial_backend)
def test_apply_runtime_state_updates_runner_attn_backend(self):
# Switching to another step config must repoint the runner backend at
# that config's draft-extend backend (read by the draft-extend forward).
new_extend_backend = object()
worker, dw = self._make_adaptive_worker(object())
state = SpecRuntimeState(
speculative_num_steps=3,
speculative_num_draft_tokens=4,
draft_attn_backend=object(),
cuda_graph_runner=object(),
target_attn_backend=object(),
target_graph_runner=object(),
draft_extend_attn_backend=new_extend_backend,
cuda_graph_runner_for_draft_extend=object(),
)
worker.apply_runtime_state(state)
self.assertIs(dw.draft_runner.attn_backend, new_extend_backend)
def test_spec_v2_attn_backends_include_draft_extend_fallback(self):
target_backend = object()
decode_backend = object()