[EAGLE] Prune draft-extend logits to selected rows (#35546)
This commit is contained in:
@@ -273,8 +273,8 @@ class LogitsMetadata:
|
|||||||
|
|
||||||
mm_input_embeds: Optional[torch.Tensor] = None
|
mm_input_embeds: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# DRAFT_EXTEND_V2: when set, lm_head runs only on these rows (see
|
# DRAFT_EXTEND_V2: when set, lm_head and LAST hidden capture use only these
|
||||||
# EagleDraftExtendInput.select_index).
|
# rows (see EagleDraftExtendInput.select_index).
|
||||||
draft_extend_select_index: Optional[torch.Tensor] = None
|
draft_extend_select_index: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -539,19 +539,37 @@ class LogitsProcessor(nn.Module):
|
|||||||
or logits_metadata.forward_mode.is_target_verify()
|
or logits_metadata.forward_mode.is_target_verify()
|
||||||
or logits_metadata.forward_mode.is_draft_extend_v2()
|
or logits_metadata.forward_mode.is_draft_extend_v2()
|
||||||
):
|
):
|
||||||
if logits_metadata.draft_extend_select_index is not None:
|
draft_extend_select_index = logits_metadata.draft_extend_select_index
|
||||||
# Only next_token_logits narrows to [bs, vocab]; the
|
if draft_extend_select_index is not None:
|
||||||
# FULL-capture hidden stays unpruned.
|
# The draft-extend graph returns LAST hidden states alongside
|
||||||
pruned_states = hidden_states[logits_metadata.draft_extend_select_index]
|
# selected logits. Build selected variants for every hidden-state
|
||||||
|
# representation; FULL capture below still uses the original
|
||||||
|
# unpruned tensors.
|
||||||
|
pruned_states = hidden_states[draft_extend_select_index]
|
||||||
|
pruned_states_before_norm = (
|
||||||
|
hidden_states_before_norm[draft_extend_select_index]
|
||||||
|
if hidden_states_before_norm is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
pruned_states = hidden_states
|
pruned_states = hidden_states
|
||||||
pruned_states_before_norm = hidden_states_before_norm
|
pruned_states_before_norm = hidden_states_before_norm
|
||||||
if aux_hidden_states is not None:
|
if aux_hidden_states is not None:
|
||||||
aux_pruned_states = (
|
if draft_extend_select_index is not None:
|
||||||
aux_hidden_states
|
aux_pruned_states = (
|
||||||
if isinstance(aux_hidden_states, torch.Tensor)
|
aux_hidden_states[draft_extend_select_index]
|
||||||
else [hidden for hidden in aux_hidden_states]
|
if isinstance(aux_hidden_states, torch.Tensor)
|
||||||
)
|
else [
|
||||||
|
hidden[draft_extend_select_index]
|
||||||
|
for hidden in aux_hidden_states
|
||||||
|
]
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
aux_pruned_states = (
|
||||||
|
aux_hidden_states
|
||||||
|
if isinstance(aux_hidden_states, torch.Tensor)
|
||||||
|
else [hidden for hidden in aux_hidden_states]
|
||||||
|
)
|
||||||
sample_indices = None
|
sample_indices = None
|
||||||
input_logprob_indices = None
|
input_logprob_indices = None
|
||||||
|
|
||||||
|
|||||||
@@ -2,7 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
from dataclasses import dataclass, fields
|
from dataclasses import dataclass, fields
|
||||||
from typing import Dict, Tuple
|
from typing import Collection, Dict, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -68,13 +68,15 @@ class ForwardInputBuffers:
|
|||||||
if buffer is not None:
|
if buffer is not None:
|
||||||
buffer.zero_()
|
buffer.zero_()
|
||||||
|
|
||||||
def share_buffers(self):
|
def share_buffers(self, *, exclude: Collection[str] = ()):
|
||||||
# disable share input buffer on npu due to accuracy issue
|
# disable share input buffer on npu due to accuracy issue
|
||||||
if is_npu():
|
if is_npu():
|
||||||
return
|
return
|
||||||
|
|
||||||
for f in fields(self):
|
for f in fields(self):
|
||||||
name = f.name
|
name = f.name
|
||||||
|
if name in exclude:
|
||||||
|
continue
|
||||||
buffer = getattr(self, name)
|
buffer = getattr(self, name)
|
||||||
|
|
||||||
if buffer is None:
|
if buffer is None:
|
||||||
|
|||||||
@@ -68,6 +68,7 @@ class EagleDraftExtendInputBuffers(ForwardInputBuffers):
|
|||||||
extend_seq_lens: torch.Tensor
|
extend_seq_lens: torch.Tensor
|
||||||
num_correct_drafts: torch.Tensor
|
num_correct_drafts: torch.Tensor
|
||||||
num_accept_tokens: torch.Tensor
|
num_accept_tokens: torch.Tensor
|
||||||
|
select_index: torch.Tensor
|
||||||
next_token_logits_buffer: torch.Tensor
|
next_token_logits_buffer: torch.Tensor
|
||||||
global_num_tokens_gpu: Optional[torch.Tensor]
|
global_num_tokens_gpu: Optional[torch.Tensor]
|
||||||
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor]
|
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor]
|
||||||
@@ -190,6 +191,11 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
num_accept_tokens = torch.full(
|
num_accept_tokens = torch.full(
|
||||||
(self.max_bs,), self.captured_req_width, dtype=torch.int32
|
(self.max_bs,), self.captured_req_width, dtype=torch.int32
|
||||||
)
|
)
|
||||||
|
select_index = (
|
||||||
|
torch.arange(self.max_bs, dtype=torch.int64) * self.captured_req_width
|
||||||
|
+ self.captured_req_width
|
||||||
|
- 1
|
||||||
|
)
|
||||||
|
|
||||||
if self.require_gathered_buffer:
|
if self.require_gathered_buffer:
|
||||||
if self.require_mlp_tp_gather:
|
if self.require_mlp_tp_gather:
|
||||||
@@ -227,7 +233,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
|
|
||||||
next_token_logits_buffer = (
|
next_token_logits_buffer = (
|
||||||
self.model_runner.graph_shared_output.get_logits_buffer(
|
self.model_runner.graph_shared_output.get_logits_buffer(
|
||||||
vocab_size, rows=self.max_bs * self.captured_req_width
|
vocab_size, rows=self.max_bs
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -258,12 +264,17 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
extend_seq_lens=extend_seq_lens,
|
extend_seq_lens=extend_seq_lens,
|
||||||
num_correct_drafts=num_correct_drafts,
|
num_correct_drafts=num_correct_drafts,
|
||||||
num_accept_tokens=num_accept_tokens,
|
num_accept_tokens=num_accept_tokens,
|
||||||
|
select_index=select_index,
|
||||||
next_token_logits_buffer=next_token_logits_buffer,
|
next_token_logits_buffer=next_token_logits_buffer,
|
||||||
global_num_tokens_gpu=global_num_tokens_gpu,
|
global_num_tokens_gpu=global_num_tokens_gpu,
|
||||||
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
|
||||||
dsa_seed_topk_capture=dsa_seed_topk_capture,
|
dsa_seed_topk_capture=dsa_seed_topk_capture,
|
||||||
)
|
)
|
||||||
self.buffers.share_buffers()
|
# The values depend on captured_req_width, while adaptive speculative
|
||||||
|
# decoding owns multiple runners whose select_index buffers all have
|
||||||
|
# shape [max_bs]. Sharing by field name and shape would alias those
|
||||||
|
# width-specific indices and can make a narrower graph gather OOB.
|
||||||
|
self.buffers.share_buffers(exclude={"select_index"})
|
||||||
|
|
||||||
self.backend = resolve_decode_backend(self)
|
self.backend = resolve_decode_backend(self)
|
||||||
|
|
||||||
@@ -343,17 +354,23 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
)
|
)
|
||||||
num_correct_drafts = buffers.num_correct_drafts[:bs]
|
num_correct_drafts = buffers.num_correct_drafts[:bs]
|
||||||
num_accept_tokens = buffers.num_accept_tokens[:bs]
|
num_accept_tokens = buffers.num_accept_tokens[:bs]
|
||||||
next_token_logits_buffer = buffers.next_token_logits_buffer[:num_tokens]
|
select_index = buffers.select_index[:bs]
|
||||||
|
next_token_logits_buffer = buffers.next_token_logits_buffer[:bs]
|
||||||
|
|
||||||
# pruned_states = num_tokens (all tokens)
|
# The worker samples only the last accepted row from each request.
|
||||||
num_tokens_for_logprob = num_tokens
|
# Keep the full tree width for the draft forward, but run the lm_head
|
||||||
|
# and logits path on those selected rows only.
|
||||||
|
num_tokens_for_logprob = bs
|
||||||
|
|
||||||
if self.require_mlp_tp_gather:
|
if self.require_mlp_tp_gather:
|
||||||
global_num_tokens_cpu = [num_tokens] * self.attn_dp_size
|
global_num_tokens_cpu = [num_tokens] * self.attn_dp_size
|
||||||
|
global_num_tokens_for_logprob_cpu = [bs] * self.attn_dp_size
|
||||||
elif self.require_attn_tp_gather:
|
elif self.require_attn_tp_gather:
|
||||||
global_num_tokens_cpu = [num_tokens]
|
global_num_tokens_cpu = [num_tokens]
|
||||||
|
global_num_tokens_for_logprob_cpu = [bs]
|
||||||
else:
|
else:
|
||||||
global_num_tokens_cpu = None
|
global_num_tokens_cpu = None
|
||||||
|
global_num_tokens_for_logprob_cpu = None
|
||||||
|
|
||||||
if global_num_tokens_cpu is not None:
|
if global_num_tokens_cpu is not None:
|
||||||
global_dp_buffer_len = sum(global_num_tokens_cpu)
|
global_dp_buffer_len = sum(global_num_tokens_cpu)
|
||||||
@@ -380,6 +397,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
num_accept_tokens=num_accept_tokens,
|
num_accept_tokens=num_accept_tokens,
|
||||||
# Padded tree width per req; drives the constant qo layout.
|
# Padded tree width per req; drives the constant qo layout.
|
||||||
num_tokens_per_req=self.captured_req_width,
|
num_tokens_per_req=self.captured_req_width,
|
||||||
|
num_tokens_for_logprob_per_req=1,
|
||||||
|
select_index=select_index,
|
||||||
)
|
)
|
||||||
|
|
||||||
forward_batch = ForwardBatch(
|
forward_batch = ForwardBatch(
|
||||||
@@ -398,6 +417,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
positions=positions,
|
positions=positions,
|
||||||
mrope_positions=mrope_positions,
|
mrope_positions=mrope_positions,
|
||||||
global_num_tokens_gpu=buffers.global_num_tokens_gpu,
|
global_num_tokens_gpu=buffers.global_num_tokens_gpu,
|
||||||
|
global_num_tokens_for_logprob_cpu=global_num_tokens_for_logprob_cpu,
|
||||||
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
|
global_num_tokens_for_logprob_gpu=buffers.global_num_tokens_for_logprob_gpu,
|
||||||
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
|
dp_padding_mode=DpPaddingMode.get_default_mode_in_cuda_graph(),
|
||||||
global_dp_buffer_len=global_dp_buffer_len,
|
global_dp_buffer_len=global_dp_buffer_len,
|
||||||
@@ -469,7 +489,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
post_warmup_hook=post_warmup_hook,
|
post_warmup_hook=post_warmup_hook,
|
||||||
)
|
)
|
||||||
|
|
||||||
def execute(self, forward_batch: ForwardBatch):
|
def execute(self, forward_batch: ForwardBatch, select_index: torch.Tensor):
|
||||||
assert forward_batch.out_cache_loc is not None
|
assert forward_batch.out_cache_loc is not None
|
||||||
self.deepep_adapter.replay()
|
self.deepep_adapter.replay()
|
||||||
buffers = self.buffers
|
buffers = self.buffers
|
||||||
@@ -526,6 +546,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
copy_srcs.append(forward_batch.spec_info.num_correct_drafts)
|
copy_srcs.append(forward_batch.spec_info.num_correct_drafts)
|
||||||
copy_dsts.append(buffers.num_accept_tokens[:raw_bs])
|
copy_dsts.append(buffers.num_accept_tokens[:raw_bs])
|
||||||
copy_srcs.append(forward_batch.spec_info.num_accept_tokens)
|
copy_srcs.append(forward_batch.spec_info.num_accept_tokens)
|
||||||
|
copy_dsts.append(buffers.select_index[:raw_bs])
|
||||||
|
copy_srcs.append(select_index)
|
||||||
_grouped_foreach_copy_(copy_dsts, copy_srcs)
|
_grouped_foreach_copy_(copy_dsts, copy_srcs)
|
||||||
|
|
||||||
# hidden_states is large + contiguous: copy_() uses the cudaMemcpyAsync
|
# hidden_states is large + contiguous: copy_() uses the cudaMemcpyAsync
|
||||||
@@ -543,9 +565,7 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
# TODO(ch-wan): support num_token_non_padded
|
# TODO(ch-wan): support num_token_non_padded
|
||||||
if self.require_gathered_buffer:
|
if self.require_gathered_buffer:
|
||||||
buffers.global_num_tokens_gpu.fill_(bs * self.captured_req_width)
|
buffers.global_num_tokens_gpu.fill_(bs * self.captured_req_width)
|
||||||
buffers.global_num_tokens_for_logprob_gpu.fill_(
|
buffers.global_num_tokens_for_logprob_gpu.fill_(bs)
|
||||||
bs * self.captured_req_width
|
|
||||||
)
|
|
||||||
|
|
||||||
if forward_batch.seq_lens_cpu is not None:
|
if forward_batch.seq_lens_cpu is not None:
|
||||||
if bs != raw_bs:
|
if bs != raw_bs:
|
||||||
@@ -607,7 +627,9 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
|
|||||||
out = self._replay_graph(shape_key, forward_batch)
|
out = self._replay_graph(shape_key, forward_batch)
|
||||||
|
|
||||||
out = LogitsProcessorOutput(
|
out = LogitsProcessorOutput(
|
||||||
next_token_logits=out.next_token_logits[:num_tokens],
|
next_token_logits=out.next_token_logits[:raw_bs],
|
||||||
hidden_states=out.hidden_states[:num_tokens],
|
# CUDA graph replay reuses its captured output storage. These states
|
||||||
|
# survive into the next draft step, so detach them from that buffer.
|
||||||
|
hidden_states=out.hidden_states[:raw_bs].clone(),
|
||||||
)
|
)
|
||||||
return out
|
return out
|
||||||
|
|||||||
@@ -315,8 +315,8 @@ class EagleDraftExtendInput(SpecInput):
|
|||||||
|
|
||||||
# Flat per-req index of each request's last accepted window row
|
# Flat per-req index of each request's last accepted window row
|
||||||
# (i * window + front + num_correct_drafts[i]). When set, the logits
|
# (i * window + front + num_correct_drafts[i]). When set, the logits
|
||||||
# processor runs lm_head only on these rows. None under gathered-buffer
|
# processor runs lm_head and LAST hidden capture only on these rows; FULL
|
||||||
# (DP) modes, whose logprob buffer sizing assumes all-row logits.
|
# hidden capture remains unpruned.
|
||||||
select_index: Optional[torch.Tensor] = None
|
select_index: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# None for draft-extend's idle batch; attention backends fall back to
|
# None for draft-extend's idle batch; attention backends fall back to
|
||||||
|
|||||||
@@ -978,7 +978,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
with canary_ctx:
|
with canary_ctx:
|
||||||
if can_run_decode_cuda_graph:
|
if can_run_decode_cuda_graph:
|
||||||
draft_logits_output = self.cuda_graph_runner_for_draft_extend.execute(
|
draft_logits_output = self.cuda_graph_runner_for_draft_extend.execute(
|
||||||
forward_batch
|
forward_batch, select_index
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
draft_logits_output = self.draft_runner.forward(
|
draft_logits_output = self.draft_runner.forward(
|
||||||
@@ -1008,15 +1008,16 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
dsa_seed_topk_indices = dsa_extend_topk_capture[select_index]
|
dsa_seed_topk_indices = dsa_extend_topk_capture[select_index]
|
||||||
|
|
||||||
# Reorganize the spec info for the next batch
|
# Reorganize the spec info for the next batch
|
||||||
draft_logits_output.next_token_logits = draft_logits_output.next_token_logits[
|
if not can_run_decode_cuda_graph:
|
||||||
select_index
|
draft_logits_output.next_token_logits = (
|
||||||
]
|
draft_logits_output.next_token_logits[select_index]
|
||||||
if draft_logits_output.hidden_states is not None:
|
)
|
||||||
draft_logits_output.hidden_states = draft_logits_output.hidden_states[
|
if draft_logits_output.hidden_states is not None:
|
||||||
select_index
|
draft_logits_output.hidden_states = draft_logits_output.hidden_states[
|
||||||
]
|
select_index
|
||||||
# The draft-extend graph only anchors full logits; selected-row topk is
|
]
|
||||||
# owned by the worker for both graph and eager paths.
|
# Selected-row top-k remains worker-owned for both graph and eager
|
||||||
|
# paths; the graph runner only moves the row selection before lm_head.
|
||||||
if get_spec().speculative_use_rejection_sampling:
|
if get_spec().speculative_use_rejection_sampling:
|
||||||
ret_draft_probs, ret_topk_p, ret_topk_index = sample_draft_proposal(
|
ret_draft_probs, ret_topk_p, ret_topk_index = sample_draft_proposal(
|
||||||
draft_logits_output.next_token_logits,
|
draft_logits_output.next_token_logits,
|
||||||
|
|||||||
+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:
|
def _assert_draft_extend_v2_outputs_close(actual, expected, settings) -> None:
|
||||||
# DRAFT_EXTEND_V2 graph runner only anchors the full-row
|
# DRAFT_EXTEND_V2 graph runner anchors selected-row next_token_logits and
|
||||||
# `next_token_logits` / `hidden_states`; the selected-row `topk_p` /
|
# hidden_states. Selected-row topk_p / topk_index remain worker-owned and
|
||||||
# `topk_index` are owned by EAGLEWorkerV2 and computed *after* replay (see
|
# are computed after replay, so compare only what the graph anchors.
|
||||||
# `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.
|
|
||||||
torch.testing.assert_close(
|
torch.testing.assert_close(
|
||||||
actual.next_token_logits,
|
actual.next_token_logits,
|
||||||
expected.next_token_logits,
|
expected.next_token_logits,
|
||||||
@@ -582,6 +579,21 @@ def _run_eagle_draft_extend_eager(
|
|||||||
return ret
|
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(
|
def run_eagle_draft_extend_cuda_graph_runner_case(
|
||||||
testcase,
|
testcase,
|
||||||
case,
|
case,
|
||||||
@@ -613,6 +625,11 @@ def run_eagle_draft_extend_cuda_graph_runner_case(
|
|||||||
settings,
|
settings,
|
||||||
)
|
)
|
||||||
expected = _run_eagle_draft_extend_eager(eager_worker, eager_batch, 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(
|
graph_fixture, graph_worker, graph_backend = _build_eagle_draft_extend_fixture(
|
||||||
testcase,
|
testcase,
|
||||||
@@ -636,7 +653,7 @@ def run_eagle_draft_extend_cuda_graph_runner_case(
|
|||||||
adapter.prepare_replay_state(graph_fixture, case, draft_inputs, settings)
|
adapter.prepare_replay_state(graph_fixture, case, draft_inputs, settings)
|
||||||
|
|
||||||
testcase.assertTrue(graph_runner.can_run_graph(graph_batch))
|
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)
|
adapter.assert_outputs_close(actual, expected, settings)
|
||||||
finally:
|
finally:
|
||||||
_reset_cuda_graph_test_buffers()
|
_reset_cuda_graph_test_buffers()
|
||||||
@@ -663,6 +680,8 @@ class _EagleDraftExtendForward(nn.Module):
|
|||||||
|
|
||||||
def _select_logits_positions(self, forward_batch: ForwardBatch) -> torch.Tensor:
|
def _select_logits_positions(self, forward_batch: ForwardBatch) -> torch.Tensor:
|
||||||
if forward_batch.forward_mode.is_draft_extend_v2():
|
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(
|
return torch.arange(
|
||||||
forward_batch.input_ids.shape[0],
|
forward_batch.input_ids.shape[0],
|
||||||
dtype=torch.int64,
|
dtype=torch.int64,
|
||||||
@@ -689,11 +708,12 @@ class _EagleDraftExtendForward(nn.Module):
|
|||||||
|
|
||||||
hidden_states = hidden_states + self.token_embed(input_ids)
|
hidden_states = hidden_states + self.token_embed(input_ids)
|
||||||
hidden_states = self.module(hidden_states, forward_batch)
|
hidden_states = self.module(hidden_states, forward_batch)
|
||||||
logits = self.lm_head(hidden_states).float()
|
|
||||||
select_index = self._select_logits_positions(forward_batch)
|
select_index = self._select_logits_positions(forward_batch)
|
||||||
|
hidden_states = hidden_states[select_index]
|
||||||
|
logits = self.lm_head(hidden_states).float()
|
||||||
return LogitsProcessorOutput(
|
return LogitsProcessorOutput(
|
||||||
next_token_logits=logits[select_index],
|
next_token_logits=logits,
|
||||||
hidden_states=hidden_states[select_index],
|
hidden_states=hidden_states,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -26,6 +26,7 @@ from sglang.srt.model_executor.cuda_graph_buffer_registry import (
|
|||||||
GraphSlot,
|
GraphSlot,
|
||||||
PaddingPolicy,
|
PaddingPolicy,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.model_executor.input_buffers import ForwardInputBuffers
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
||||||
@@ -59,6 +60,12 @@ class _MiniForwardBatch:
|
|||||||
spec_info: Optional[object] = None
|
spec_info: Optional[object] = None
|
||||||
|
|
||||||
|
|
||||||
|
@dataclasses.dataclass
|
||||||
|
class _PoolInputBuffers(ForwardInputBuffers):
|
||||||
|
input_ids: torch.Tensor
|
||||||
|
select_index: torch.Tensor
|
||||||
|
|
||||||
|
|
||||||
def _make_registry(max_bs: int = 8, max_num_tokens: int = 16):
|
def _make_registry(max_bs: int = 8, max_num_tokens: int = 16):
|
||||||
return CudaGraphBufferRegistry(
|
return CudaGraphBufferRegistry(
|
||||||
device=torch.device("cpu"),
|
device=torch.device("cpu"),
|
||||||
@@ -642,6 +649,27 @@ class TestPoolBackedAlloc(unittest.TestCase):
|
|||||||
small_first, big_after = _ptrs(16, 32)
|
small_first, big_after = _ptrs(16, 32)
|
||||||
self.assertNotEqual(small_first.data_ptr(), big_after.data_ptr())
|
self.assertNotEqual(small_first.data_ptr(), big_after.data_ptr())
|
||||||
|
|
||||||
|
def test_forward_input_buffers_can_exclude_width_specific_fields(self):
|
||||||
|
first = _PoolInputBuffers(
|
||||||
|
input_ids=torch.zeros(4, dtype=torch.int64),
|
||||||
|
select_index=torch.tensor([1, 3], dtype=torch.int64),
|
||||||
|
)
|
||||||
|
second = _PoolInputBuffers(
|
||||||
|
input_ids=torch.ones(4, dtype=torch.int64),
|
||||||
|
select_index=torch.tensor([3, 7], dtype=torch.int64),
|
||||||
|
)
|
||||||
|
|
||||||
|
first.share_buffers()
|
||||||
|
second.share_buffers(exclude={"select_index"})
|
||||||
|
|
||||||
|
self.assertEqual(first.input_ids.data_ptr(), second.input_ids.data_ptr())
|
||||||
|
self.assertNotEqual(
|
||||||
|
first.select_index.data_ptr(), second.select_index.data_ptr()
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(
|
||||||
|
second.select_index, torch.tensor([3, 7], dtype=torch.int64)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TestBuildDecodeRegistry(unittest.TestCase):
|
class TestBuildDecodeRegistry(unittest.TestCase):
|
||||||
"""``build_decode_registry`` registers the always-on FB-shared decode
|
"""``build_decode_registry`` registers the always-on FB-shared decode
|
||||||
|
|||||||
@@ -0,0 +1,98 @@
|
|||||||
|
import unittest
|
||||||
|
|
||||||
|
import torch
|
||||||
|
|
||||||
|
from sglang.srt.layers.aux_hidden_states import pack_aux_hidden_states
|
||||||
|
from sglang.srt.layers.logits_processor import LogitsMetadata, LogitsProcessor
|
||||||
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
|
CaptureHiddenMode,
|
||||||
|
ForwardMode,
|
||||||
|
)
|
||||||
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
|
register_cpu_ci(est_time=1, suite="base-a-test-cpu")
|
||||||
|
|
||||||
|
|
||||||
|
class TestEagleDraftExtendLogitsPruning(unittest.TestCase):
|
||||||
|
def setUp(self):
|
||||||
|
self.hidden_states = torch.arange(12, dtype=torch.float32).reshape(6, 2)
|
||||||
|
self.select_index = torch.tensor([2, 4], dtype=torch.int64)
|
||||||
|
|
||||||
|
def _metadata(self, capture_hidden_mode=CaptureHiddenMode.LAST):
|
||||||
|
return LogitsMetadata(
|
||||||
|
forward_mode=ForwardMode.DRAFT_EXTEND_V2,
|
||||||
|
capture_hidden_mode=capture_hidden_mode,
|
||||||
|
draft_extend_select_index=self.select_index,
|
||||||
|
)
|
||||||
|
|
||||||
|
def _stored_hidden_states(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
hidden_states_before_norm=None,
|
||||||
|
aux_hidden_states=None,
|
||||||
|
capture_hidden_mode=CaptureHiddenMode.LAST,
|
||||||
|
):
|
||||||
|
metadata = self._metadata(capture_hidden_mode)
|
||||||
|
(
|
||||||
|
pruned_states,
|
||||||
|
pruned_states_before_norm,
|
||||||
|
aux_pruned_states,
|
||||||
|
sample_indices,
|
||||||
|
_,
|
||||||
|
_,
|
||||||
|
) = LogitsProcessor._get_pruned_states(
|
||||||
|
None,
|
||||||
|
self.hidden_states,
|
||||||
|
hidden_states_before_norm,
|
||||||
|
aux_hidden_states,
|
||||||
|
metadata,
|
||||||
|
)
|
||||||
|
return LogitsProcessor._get_hidden_states_to_store(
|
||||||
|
None,
|
||||||
|
self.hidden_states,
|
||||||
|
hidden_states_before_norm,
|
||||||
|
aux_hidden_states,
|
||||||
|
pruned_states,
|
||||||
|
pruned_states_before_norm,
|
||||||
|
aux_pruned_states,
|
||||||
|
sample_indices,
|
||||||
|
metadata,
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_last_hidden_states_use_selected_rows(self):
|
||||||
|
actual = self._stored_hidden_states()
|
||||||
|
torch.testing.assert_close(actual, self.hidden_states[self.select_index])
|
||||||
|
|
||||||
|
def test_last_pre_norm_hidden_states_use_selected_rows(self):
|
||||||
|
hidden_states_before_norm = self.hidden_states + 100
|
||||||
|
actual = self._stored_hidden_states(
|
||||||
|
hidden_states_before_norm=hidden_states_before_norm
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(actual, hidden_states_before_norm[self.select_index])
|
||||||
|
|
||||||
|
def test_last_aux_hidden_states_use_selected_rows(self):
|
||||||
|
aux_hidden_states = [self.hidden_states + 100, self.hidden_states + 200]
|
||||||
|
actual = self._stored_hidden_states(aux_hidden_states=aux_hidden_states)
|
||||||
|
expected = pack_aux_hidden_states(
|
||||||
|
[hidden[self.select_index] for hidden in aux_hidden_states]
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(actual, expected)
|
||||||
|
|
||||||
|
def test_last_packed_aux_hidden_states_use_selected_rows(self):
|
||||||
|
aux_hidden_states = torch.cat(
|
||||||
|
[self.hidden_states + 100, self.hidden_states + 200], dim=-1
|
||||||
|
)
|
||||||
|
actual = self._stored_hidden_states(aux_hidden_states=aux_hidden_states)
|
||||||
|
torch.testing.assert_close(actual, aux_hidden_states[self.select_index])
|
||||||
|
|
||||||
|
def test_full_hidden_capture_stays_unpruned(self):
|
||||||
|
aux_hidden_states = [self.hidden_states + 100, self.hidden_states + 200]
|
||||||
|
actual = self._stored_hidden_states(
|
||||||
|
aux_hidden_states=aux_hidden_states,
|
||||||
|
capture_hidden_mode=CaptureHiddenMode.FULL,
|
||||||
|
)
|
||||||
|
torch.testing.assert_close(actual, pack_aux_hidden_states(aux_hidden_states))
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
unittest.main()
|
||||||
Reference in New Issue
Block a user