From 8f337682bda92aa739f6d83978eae725535b386a Mon Sep 17 00:00:00 2001 From: Tarushii Goel Date: Tue, 7 Apr 2026 04:38:04 +0800 Subject: [PATCH] [sgl] potential chained spec v2 fixes (#22041) Co-authored-by: Mook Co-authored-by: yudian0504 --- .../multi_layer_eagle_draft_extend_cuda_graph_runner.py | 2 +- .../srt/speculative/multi_layer_eagle_worker_v2.py | 9 +++++++++ 2 files changed, 10 insertions(+), 1 deletion(-) diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 5fd43db24..4eb3d55c4 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -437,7 +437,7 @@ class MultiLayerEagleDraftExtendCudaGraphRunner: self.eagle_worker.chain_mtp_hidden_states and ret.hidden_states is not None ): - self.hidden_states[:num_tokens].copy_(ret.hidden_states[:num_tokens]) + buffers.hidden_states[:num_tokens].copy_(ret.hidden_states[:num_tokens]) select_index = ( torch.arange(bs, device=self.model_runner.device) diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index fbad4adb9..65bfb048e 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -506,6 +506,15 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker): dim=-1, ) ret_topk_p, ret_topk_index = fast_topk(probs, self.topk, dim=-1) + # Chain-style: use this step's output hidden_states as next step's input + if ( + self.chain_mtp_hidden_states + and step < self.speculative_num_steps - 1 + and draft_logits_output.logits_output.hidden_states is not None + ): + forward_batch.spec_info.hidden_states = ( + draft_logits_output.logits_output.hidden_states + ) if forward_batch.extend_seq_lens is not None: rotate_input_ids_triton( forward_batch.input_ids,