Fix Eagle draft decode positions (#25015)

This commit is contained in:
Lianmin Zheng
2026-05-12 14:04:58 -07:00
committed by GitHub
parent 486b54714b
commit 48fc26a814
3 changed files with 4 additions and 3 deletions
@@ -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)
@@ -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
@@ -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(