[EAGLE] Prune draft-extend logits to selected rows (#35546)
This commit is contained in:
+30
-10
@@ -388,12 +388,9 @@ def run_mla_draft_extend_v2_cuda_graph_case(
|
||||
|
||||
|
||||
def _assert_draft_extend_v2_outputs_close(actual, expected, settings) -> None:
|
||||
# DRAFT_EXTEND_V2 graph runner only anchors the full-row
|
||||
# `next_token_logits` / `hidden_states`; the selected-row `topk_p` /
|
||||
# `topk_index` are owned by EAGLEWorkerV2 and computed *after* replay (see
|
||||
# `eagle_worker_v2._draft_extend_for_decode`). The graph computes no topk
|
||||
# and `EAGLEDraftExtendCudaGraphRunner.replay` returns no topk fields, so
|
||||
# the runner-mode reference must only compare what the graph anchors.
|
||||
# DRAFT_EXTEND_V2 graph runner anchors selected-row next_token_logits and
|
||||
# hidden_states. Selected-row topk_p / topk_index remain worker-owned and
|
||||
# are computed after replay, so compare only what the graph anchors.
|
||||
torch.testing.assert_close(
|
||||
actual.next_token_logits,
|
||||
expected.next_token_logits,
|
||||
@@ -582,6 +579,21 @@ def _run_eagle_draft_extend_eager(
|
||||
return ret
|
||||
|
||||
|
||||
def _draft_extend_select_index(
|
||||
batch: ForwardBatch, settings: EagleDraftRunnerSettings
|
||||
) -> torch.Tensor:
|
||||
return (
|
||||
torch.arange(
|
||||
batch.batch_size,
|
||||
dtype=torch.int64,
|
||||
device=batch.input_ids.device,
|
||||
)
|
||||
* settings.speculative_num_draft_tokens
|
||||
+ batch.spec_info.num_accept_tokens
|
||||
- 1
|
||||
)
|
||||
|
||||
|
||||
def run_eagle_draft_extend_cuda_graph_runner_case(
|
||||
testcase,
|
||||
case,
|
||||
@@ -613,6 +625,11 @@ def run_eagle_draft_extend_cuda_graph_runner_case(
|
||||
settings,
|
||||
)
|
||||
expected = _run_eagle_draft_extend_eager(eager_worker, eager_batch, settings)
|
||||
select_index = _draft_extend_select_index(eager_batch, settings)
|
||||
expected = LogitsProcessorOutput(
|
||||
next_token_logits=expected.next_token_logits[select_index],
|
||||
hidden_states=expected.hidden_states[select_index],
|
||||
)
|
||||
|
||||
graph_fixture, graph_worker, graph_backend = _build_eagle_draft_extend_fixture(
|
||||
testcase,
|
||||
@@ -636,7 +653,7 @@ def run_eagle_draft_extend_cuda_graph_runner_case(
|
||||
adapter.prepare_replay_state(graph_fixture, case, draft_inputs, settings)
|
||||
|
||||
testcase.assertTrue(graph_runner.can_run_graph(graph_batch))
|
||||
actual = graph_runner.execute(graph_batch)
|
||||
actual = graph_runner.execute(graph_batch, select_index)
|
||||
adapter.assert_outputs_close(actual, expected, settings)
|
||||
finally:
|
||||
_reset_cuda_graph_test_buffers()
|
||||
@@ -663,6 +680,8 @@ class _EagleDraftExtendForward(nn.Module):
|
||||
|
||||
def _select_logits_positions(self, forward_batch: ForwardBatch) -> torch.Tensor:
|
||||
if forward_batch.forward_mode.is_draft_extend_v2():
|
||||
if forward_batch.spec_info.select_index is not None:
|
||||
return forward_batch.spec_info.select_index
|
||||
return torch.arange(
|
||||
forward_batch.input_ids.shape[0],
|
||||
dtype=torch.int64,
|
||||
@@ -689,11 +708,12 @@ class _EagleDraftExtendForward(nn.Module):
|
||||
|
||||
hidden_states = hidden_states + self.token_embed(input_ids)
|
||||
hidden_states = self.module(hidden_states, forward_batch)
|
||||
logits = self.lm_head(hidden_states).float()
|
||||
select_index = self._select_logits_positions(forward_batch)
|
||||
hidden_states = hidden_states[select_index]
|
||||
logits = self.lm_head(hidden_states).float()
|
||||
return LogitsProcessorOutput(
|
||||
next_token_logits=logits[select_index],
|
||||
hidden_states=hidden_states[select_index],
|
||||
next_token_logits=logits,
|
||||
hidden_states=hidden_states,
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user