[Spec] Fix return_hidden_states under spec V2 (issue #26163) (#28496)

Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
This commit is contained in:
Khoa Pham
2026-06-17 16:21:36 -07:00
committed by GitHub
co-authored by Claude Opus 4.7
parent b1d18d562b
commit e4fd613def
3 changed files with 92 additions and 6 deletions
@@ -668,9 +668,21 @@ class SchedulerBatchResultProcessor:
)
if req.return_hidden_states and logits_output.hidden_states is not None:
req.hidden_states.append(
logits_output.hidden_states[i].cpu().clone().tolist()
)
if batch.spec_algorithm.is_none():
req.hidden_states.append(
logits_output.hidden_states[i].cpu().clone().tolist()
)
else:
# Spec V2: hidden_states is [bs * speculative_num_draft_tokens, hidden_dim].
stride = result.speculative_num_draft_tokens
accept_len = result.num_correct_drafts_per_req_cpu[i] + 1
start = i * stride
req.hidden_states.extend(
logits_output.hidden_states[start : start + accept_len]
.cpu()
.clone()
.tolist()
)
if req.grammar is not None:
self._apply_decode_grammar(
@@ -474,9 +474,14 @@ class _GenerationStreamAccumulator:
self.output_token_ids_logprobs_idx.append([])
if self.return_hidden_states:
self.output_hidden_states.append(
req.hidden_states if req.return_hidden_states else None
)
if req.return_hidden_states:
# Mirror output_ids_through_stop: spec verify steps can overshoot finished_len.
hs = req.hidden_states
if req.finished_len is not None:
hs = hs[: req.finished_len]
self.output_hidden_states.append(hs)
else:
self.output_hidden_states.append(None)
if self.return_routed_experts:
self.routed_experts.append(
req.routed_experts if req.return_routed_experts else None