From 850887dc636dc3fb44b6708112953bdbef2191db Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Sun, 24 May 2026 18:49:33 -0700 Subject: [PATCH] [Spec] fix EAGLE v2 verify metadata init order on non-cuda-graph path (#26244) --- python/sglang/srt/speculative/eagle_info_v2.py | 8 +++----- python/sglang/srt/speculative/eagle_worker_v2.py | 6 ++++-- 2 files changed, 7 insertions(+), 7 deletions(-) diff --git a/python/sglang/srt/speculative/eagle_info_v2.py b/python/sglang/srt/speculative/eagle_info_v2.py index a7c507297..dd426f2e8 100644 --- a/python/sglang/srt/speculative/eagle_info_v2.py +++ b/python/sglang/srt/speculative/eagle_info_v2.py @@ -306,11 +306,9 @@ class EagleVerifyInputV2Mixin: ) if can_run_cuda_graph: target_worker.model_runner.graph_runner.replay_prepare(verify_forward_batch) - else: - if not batch.forward_mode.is_idle(): - target_worker.model_runner.attn_backend.init_forward_metadata( - verify_forward_batch - ) + # Non-cuda-graph: defer init to forward_extend, which runs after + # `_forward_raw -> prepare_mlp_sync_batch` pads the batch. Initing + # here would use pre-pad shapes and trip DSv4 indexer shape match. return verify_forward_batch, can_run_cuda_graph diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 18a947e1a..eb5c40a77 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -1039,12 +1039,14 @@ class EAGLEWorkerV2(BaseSpecWorker): verify_input.retrieve_next_token.shape ).cpu() - # Run target verify batch in the main compute stream (GPU compute) + # Run target verify batch in the main compute stream (GPU compute). + # Only skip metadata init when cuda-graph already ran replay_prepare; + # the non-cuda-graph path needs forward_extend's init (post-pad). forward_batch_output = self.target_worker.forward_batch_generation( batch=None, forward_batch=verify_forward_batch, is_verify=True, - skip_attn_backend_init=True, + skip_attn_backend_init=can_run_cuda_graph, ) logits_output = forward_batch_output.logits_output