From d601edab73b2225eb153998b1534a306876f9216 Mon Sep 17 00:00:00 2001 From: Cheng Wan <54331508+ch-wan@users.noreply.github.com> Date: Fri, 12 Jun 2026 15:47:05 -0700 Subject: [PATCH] [Spec] Fix EagleDraftWorker draft-extend attn backend assignment (#28096) Co-authored-by: Claude Opus 4.8 (1M context) --- .../sglang/srt/speculative/eagle_worker_v2.py | 10 +++ .../test_eagle_worker_v2_topk1_fastpath.py | 72 +++++++++++++++++++ 2 files changed, 82 insertions(+) diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index bd22c68fa..03b378faa 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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, diff --git a/test/registered/unit/spec/test_eagle_worker_v2_topk1_fastpath.py b/test/registered/unit/spec/test_eagle_worker_v2_topk1_fastpath.py index e94732d46..864b54ef1 100644 --- a/test/registered/unit/spec/test_eagle_worker_v2_topk1_fastpath.py +++ b/test/registered/unit/spec/test_eagle_worker_v2_topk1_fastpath.py @@ -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()