From 48fc26a814b8a71ffcb1e6a66b2a57393fd75e04 Mon Sep 17 00:00:00 2001 From: Lianmin Zheng Date: Tue, 12 May 2026 14:04:58 -0700 Subject: [PATCH] Fix Eagle draft decode positions (#25015) --- python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py | 3 ++- python/sglang/srt/speculative/eagle_worker.py | 2 +- python/sglang/srt/speculative/eagle_worker_v2.py | 2 +- 3 files changed, 4 insertions(+), 3 deletions(-) diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index b0cd9a32d..1dd4bb039 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -364,7 +364,7 @@ class EAGLEDraftCudaGraphRunner: ) set_is_extend_in_batch(False) - # Backup two fields, which will be modified in-place in `draft_forward`. + # Backup fields that are modified in-place in `draft_forward`. output_cache_loc_backup = forward_batch.out_cache_loc hidden_states_backup = forward_batch.spec_info.hidden_states @@ -372,6 +372,7 @@ class EAGLEDraftCudaGraphRunner: forward_batch.out_cache_loc = output_cache_loc_backup forward_batch.spec_info.hidden_states = hidden_states_backup + forward_batch.positions.sub_(self.eagle_worker.speculative_num_steps - 1) return ret self.deepep_adapter.capture(is_extend_in_batch=False) diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index f4bd641f4..cf6339074 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -893,7 +893,6 @@ class EAGLEWorker(TpModelWorker): ): out_cache_loc = out_cache_loc.contiguous() forward_batch.out_cache_loc = out_cache_loc[i] - forward_batch.positions.add_(1) forward_batch.attn_backend = self.draft_attn_backend.attn_backends[i] spec_info.hidden_states = hidden_states @@ -913,6 +912,7 @@ class EAGLEWorker(TpModelWorker): if self.hot_token_id is not None: topk_index = self.hot_token_id[topk_index] hidden_states = logits_output.hidden_states + forward_batch.positions.add_(1) parent_list, top_scores_index, draft_tokens = organize_draft_results( score_list, token_list, parents_list, self.speculative_num_draft_tokens diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 4478ddf6d..2403aab4c 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -457,7 +457,6 @@ class EagleDraftWorker(BaseDraftWorker): # Set inputs forward_batch.input_ids = input_ids forward_batch.out_cache_loc = out_cache_loc[i] - forward_batch.positions.add_(1) forward_batch.attn_backend = self.draft_attn_backend.attn_backends[i] spec_info.hidden_states = hidden_states @@ -477,6 +476,7 @@ class EagleDraftWorker(BaseDraftWorker): if self.hot_token_id is not None: topk_index = self.hot_token_id[topk_index] hidden_states = logits_output.hidden_states + forward_batch.positions.add_(1) # Organize the results score_list = torch.cat(score_list, dim=1).flatten(