From ac99794e64e054224dbdd52b4318629ad06540c8 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Wed, 3 Jun 2026 15:40:05 -0400 Subject: [PATCH] Reland spec v2 tree drafting (eagle topk>1) with page_size==1 (#26866) (#26997) Co-authored-by: Alison Shao <54658187+alisonshao@users.noreply.github.com> --- .../sglang/srt/arg_groups/speculative_hook.py | 14 ++++- .../attention/flashattention_backend.py | 14 +++++ .../srt/layers/attention/triton_backend.py | 10 ++- python/sglang/srt/mem_cache/memory_pool.py | 35 +++++++++++ .../srt/model_executor/forward_batch_info.py | 9 ++- .../model_runner_kv_cache_mixin.py | 6 ++ .../eagle_draft_extend_cuda_graph_runner.py | 4 +- .../sglang/srt/speculative/eagle_worker_v2.py | 63 +++++++++++++++++-- .../multi_layer_eagle_worker_v2.py | 3 +- .../srt/speculative/triton_ops/eagle.py | 5 +- .../spec/eagle/test_spec_eagle_stress.py | 16 ++++- .../spec/eagle/test_spec_eagle_topk.py | 15 ++++- 12 files changed, 176 insertions(+), 18 deletions(-) diff --git a/python/sglang/srt/arg_groups/speculative_hook.py b/python/sglang/srt/arg_groups/speculative_hook.py index 71008d6cd..8d61095a5 100644 --- a/python/sglang/srt/arg_groups/speculative_hook.py +++ b/python/sglang/srt/arg_groups/speculative_hook.py @@ -277,13 +277,25 @@ def _handle_eagle_family(server_args: "ServerArgs") -> None: ) spec_v1_reason = None + # mamba / linear-attn state models only support topk == 1 on spec v2. + # mamba2_cache_params exists iff the config carries such state; check the + # class descriptor so the property getter is not invoked. + text_config = server_args.get_model_config().hf_config.get_text_config() + is_mamba_state_model = hasattr(type(text_config), "mamba2_cache_params") if ( server_args.speculative_eagle_topk is not None and server_args.speculative_eagle_topk > 1 + and (server_args.page_size > 1 or is_mamba_state_model) and not server_args.disable_overlap_schedule ): + # Spec v2 topk > 1 only supports page_size == 1 on non-mamba models; + # page_size > 1 (partial-page dup) isn't ported to v2 yet -> fall back to v1. server_args.disable_overlap_schedule = True - spec_v1_reason = "spec v2 currently only supports topk = 1" + spec_v1_reason = ( + "spec v2 topk > 1 is not supported for mamba/linear-attn models" + if is_mamba_state_model + else "spec v2 topk > 1 currently requires page_size == 1" + ) elif ( not envs.SGLANG_ENABLE_SPEC_V2.get() and not server_args.disable_overlap_schedule diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index a5b637c1f..351dc90ee 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -398,6 +398,20 @@ class FlashAttentionBackend(AttentionBackend): metadata.page_table = self.req_to_token_pool.req_to_token[ forward_batch.req_pool_indices, : metadata.max_seq_len_k ] + elif self.speculative_num_steps == 0: + # Draft-extend's idle batch (padded for DP MLP-sync) has no + # tree; build plain metadata (padded output is discarded). + metadata.cache_seqlens_int32 = seqlens_in_batch.to(torch.int32) + metadata.max_seq_len_k = forward_batch.seq_lens_cpu.max().item() + metadata.cu_seqlens_q = torch.arange( + 0, batch_size + 1, dtype=torch.int32, device=device + ) + metadata.cu_seqlens_k = torch.nn.functional.pad( + torch.cumsum(seqlens_in_batch, dim=0, dtype=torch.int32), (1, 0) + ) + metadata.page_table = self.req_to_token_pool.req_to_token[ + forward_batch.req_pool_indices, : metadata.max_seq_len_k + ] else: metadata.cache_seqlens_int32 = (seqlens_in_batch).to(torch.int32) metadata.max_seq_len_q = self.topk diff --git a/python/sglang/srt/layers/attention/triton_backend.py b/python/sglang/srt/layers/attention/triton_backend.py index f280f8b9f..4af1724aa 100644 --- a/python/sglang/srt/layers/attention/triton_backend.py +++ b/python/sglang/srt/layers/attention/triton_backend.py @@ -534,9 +534,15 @@ class TritonAttnBackend(AttentionBackend): spec_info = forward_batch.spec_info if forward_batch.forward_mode.is_decode_or_idle(): - if spec_info is None: + if spec_info is None or spec_info.kv_indptr is None: + # kv_indptr is None for draft-extend's idle batch (no tree + # indices); build plain metadata from seq_lens. + # gpu_only: seq_lens_sum may be None; ub-allocate is safe (ragged write). + seq_lens_sum = forward_batch.seq_lens_sum + if seq_lens_sum is None: + seq_lens_sum = bs * self.max_context_len kv_indices = torch.empty( - forward_batch.seq_lens_sum, dtype=torch.int64, device=self.device + seq_lens_sum, dtype=torch.int64, device=self.device ) kv_indptr = self._fill_kv_indptr_and_indices( bs, diff --git a/python/sglang/srt/mem_cache/memory_pool.py b/python/sglang/srt/mem_cache/memory_pool.py index b0105d7cf..bee9ae6bd 100644 --- a/python/sglang/srt/mem_cache/memory_pool.py +++ b/python/sglang/srt/mem_cache/memory_pool.py @@ -854,6 +854,11 @@ class MHATokenToKVPool(KVCache): self.same_kv_dim = self.head_dim == self.v_head_dim def _init_kv_copy_and_warmup(self): + # Zero-layer pool (e.g. all-SWA model's full sub-pool) has no buffers. + if self.layer_num == 0: + self._kv_copy_config = None + return + # Heuristics for KV copy tiling _KV_COPY_STRIDE_THRESHOLD_LARGE = 8192 _KV_COPY_STRIDE_THRESHOLD_MEDIUM = 4096 @@ -1090,6 +1095,10 @@ class MHATokenToKVPool(KVCache): ) def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor): + # Zero-layer pool (e.g. all-SWA model's full sub-pool) has no buffers. + if self.layer_num == 0: + return + # Catch stale indices here instead of as illegal-addr or silent KV corruption. size_limit = self.size + self.page_size maybe_detect_oob(tgt_loc, 0, size_limit, "move_kv_cache tgt_loc") @@ -1828,6 +1837,20 @@ class MLATokenToKVPool(KVCache): get_mla_kv_buffer_triton(kv_buffer, loc, cache_k_nope, cache_k_rope) return cache_k_nope, cache_k_rope + def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor): + """Relocate accepted-token combined MLA KV (latent + rope) per layer.""" + size_limit = self.size + self.page_size + maybe_detect_oob(tgt_loc, 0, size_limit, "move_kv_cache tgt_loc") + maybe_detect_oob(src_loc, 0, size_limit, "move_kv_cache src_loc") + + if tgt_loc.numel() == 0: + return + + tgt_loc_flat = tgt_loc.view(-1).long() + src_loc_flat = src_loc.view(-1).long() + for kv_cache in self.kv_buffer: + kv_cache[tgt_loc_flat] = kv_cache[src_loc_flat] + def get_cpu_copy(self, indices, mamba_indices=None): current_platform.synchronize() kv_cache_cpu = [] @@ -2073,6 +2096,18 @@ class DSATokenToKVPool(MLATokenToKVPool): del self.kv_buffer del self.index_k_with_scale_buffer + def move_kv_cache(self, tgt_loc: torch.Tensor, src_loc: torch.Tensor): + """Move latent KV and the DSA indexer cache (key + scale) in lockstep.""" + super().move_kv_cache(tgt_loc, src_loc) + + if tgt_loc.numel() == 0: + return + + tgt_loc_flat = tgt_loc.view(-1).long() + src_loc_flat = src_loc.view(-1).long() + for index_k in self.index_k_with_scale_buffer: + index_k[tgt_loc_flat] = index_k[src_loc_flat] + def get_index_k_with_scale_buffer(self, layer_id: int) -> torch.Tensor: if self.layer_transfer_counter is not None: self.layer_transfer_counter.wait_until(layer_id - self.start_layer) diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 8a9109e42..812196ffd 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -1086,9 +1086,12 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): seq_len_fill_value = ( model_runner.attn_backend.get_cuda_graph_seq_len_fill_value() ) - self.seq_lens_sum = self.seq_lens_sum + seq_len_fill_value * ( - bs - self.seq_lens.shape[0] - ) + # Keep gpu_only batches sync-free: leave seq_lens_sum None and let the + # attention backend over-allocate from an upper bound (see #26738). + if self.seq_lens_sum is not None: + self.seq_lens_sum = self.seq_lens_sum + seq_len_fill_value * ( + bs - self.seq_lens.shape[0] + ) self.seq_lens = self._pad_tensor_to_size( self.seq_lens, bs, value=seq_len_fill_value ) diff --git a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py index 5ee04b2ad..b72252dec 100644 --- a/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py +++ b/python/sglang/srt/model_executor/model_runner_kv_cache_mixin.py @@ -499,6 +499,9 @@ class ModelRunnerKVCacheMixin: enable_kvcache_transpose=False, device=self.device, token_to_kv_pool_class=NPUMHATokenToKVPool, + enable_kv_cache_copy=( + self.server_args.speculative_algorithm is not None + ), **kwargs, ) elif self.use_mla_backend: @@ -621,6 +624,9 @@ class ModelRunnerKVCacheMixin: full_attention_layer_ids=self.model_config.full_attention_layer_ids, enable_kvcache_transpose=False, device=self.device, + enable_kv_cache_copy=( + self.server_args.speculative_algorithm is not None + ), **kwargs, ) elif config := self.mambaish_config: diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index 7f45ad648..4c7ee8cd7 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -108,7 +108,9 @@ class EAGLEDraftExtendCudaGraphRunner: self.padded_static_len = -1 # Attention backend - self.num_tokens_per_bs = self.speculative_num_steps + 1 + # Size cuda-graph buffers by num_draft_tokens (full tree width), not + # num_steps + 1, or topk > 1 draft-extend overflows them. + self.num_tokens_per_bs = model_runner.server_args.speculative_num_draft_tokens self.max_bs = max(self.capture_bs) self.max_num_token = self.max_bs * self.num_tokens_per_bs diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 4f367a256..798bf9509 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -693,8 +693,10 @@ class EagleDraftWorker(BaseDraftWorker): # Batch 2: Draft extend draft_input = EagleDraftInput( hidden_states=batch_result.logits_output.hidden_states, - num_tokens_per_req=self.speculative_num_steps + 1, - num_tokens_for_logprob_per_req=self.speculative_num_steps + 1, + # Draft-extend fills the whole tree width (num_draft_tokens) per req, + # not num_steps + 1, so DP MLP-sync padding stays consistent for topk > 1. + num_tokens_per_req=self.speculative_num_draft_tokens, + num_tokens_for_logprob_per_req=self.speculative_num_draft_tokens, ) select_index = ( torch.arange(len(batch.seq_lens), device=self.device) @@ -1232,11 +1234,13 @@ class EAGLEWorkerV2(BaseSpecWorker): if not batch.forward_mode.is_idle(): accept_tokens = predict[accept_index] bonus_tokens = torch.empty_like(accept_lens, dtype=torch.int32) + # stride = accept_tokens per-req width = accept_index.shape[1] + # (spec_steps + 1); NOT num_draft_tokens, wrong for topk > 1 trees. fill_bonus_tokens[(bs,)]( accept_tokens, accept_lens, bonus_tokens, - self.speculative_num_draft_tokens, + accept_index.shape[1], ) else: bonus_tokens = torch.empty((0,), device=self.device, dtype=torch.int32) @@ -1246,6 +1250,13 @@ class EAGLEWorkerV2(BaseSpecWorker): batch, logits_output, predict, accept_index, self.speculative_num_steps ) + if not batch.forward_mode.is_idle() and self.topk > 1: + # topk == 1 needs nothing here: the accepted path is already the front + # chain, so the whole compaction is an identity transform. + predict = self._finalize_accepted_tree_path( + batch, accept_index, accept_lens, predict, logits_output, bs + ) + next_draft_input = EagleDraftInput(bonus_tokens=bonus_tokens) # verify_forward_batch transitively holds verify-time GPU tensors @@ -1327,6 +1338,30 @@ class EAGLEWorkerV2(BaseSpecWorker): model=self.target_worker.model_runner.model, ) + def _finalize_accepted_tree_path( + self, + batch: ScheduleBatch, + accept_index: torch.Tensor, + accept_lens: torch.Tensor, + predict: torch.Tensor, + logits_output, + bs: int, + ) -> torch.Tensor: + """Tree drafting (topk > 1): move the accepted path -- KV slots, predict, + hidden_states -- to the contiguous front of each per-req block, which the + downstream chain-layout code (draft-extend select_index, committed-KV reads) + assumes. Returns compacted predict; mutates logits_output.hidden_states + (moved only when present).""" + self.move_accepted_tokens_to_target_kvcache( + batch, accept_index, accept_lens - 1 + ) + predict = self._compact_accepted_to_front(predict, accept_index, bs) + if logits_output.hidden_states is not None: + logits_output.hidden_states = self._compact_accepted_to_front( + logits_output.hidden_states, accept_index, bs + ) + return predict + def move_accepted_tokens_to_target_kvcache( self, batch: ScheduleBatch, @@ -1343,7 +1378,9 @@ class EAGLEWorkerV2(BaseSpecWorker): seq_lens is advanced by ``num_correct_drafts + 1`` to cover the bonus slot. """ bs = len(batch.seq_lens) - size = bs * self.speculative_num_draft_tokens + # accept_index element count, NOT bs * num_draft_tokens: for topk > 1 the + # tree exceeds the accepted chain, over-reading accept_index (illegal memory). + size = bs * accept_index.shape[1] # fill_accepted_out_cache_loc reads out_cache_loc[accept_index]; -1 sentinel ok. maybe_detect_oob( @@ -1380,6 +1417,24 @@ class EAGLEWorkerV2(BaseSpecWorker): tgt_cache_loc, accepted_out_cache_loc ) + def _compact_accepted_to_front( + self, x: torch.Tensor, accept_index: torch.Tensor, bs: int + ) -> torch.Tensor: + """Gather the accepted tree path to the front of each per-req block. + + ``x`` is node-indexed over the whole tree (``[bs * num_draft_tokens, ...]``), + ``accept_index`` is ``[bs, spec_steps + 1]`` global node indices (-1 padded). + Padded entries clamp to node 0 but land past accept_lens (never read); + trailing unaccepted slots stay and are freed as overshoot. + """ + nd = self.speculative_num_draft_tokens + s1 = accept_index.shape[1] # spec_steps + 1 + safe = accept_index.to(torch.int64).clamp(min=0).reshape(-1) + gathered = x[safe] + out = x.clone() + out.view(bs, nd, *x.shape[1:])[:, :s1] = gathered.view(bs, s1, *x.shape[1:]) + return out + def update_weights_from_disk(self, recv_req: UpdateWeightFromDiskReqInput): success, message = self._draft_worker.draft_runner.update_weights_from_disk( recv_req.model_path, 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 725142669..4b74c65aa 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -790,11 +790,12 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): if not batch.forward_mode.is_idle(): accept_tokens = predict[accept_index] bonus_tokens = torch.empty_like(accept_lens, dtype=torch.int32) + # stride = accept_tokens per-req width = accept_index.shape[1]. fill_bonus_tokens[(bs,)]( accept_tokens, accept_lens, bonus_tokens, - self.speculative_num_draft_tokens, + accept_index.shape[1], ) else: bonus_tokens = torch.empty((0,), device=self.device, dtype=torch.int32) diff --git a/python/sglang/srt/speculative/triton_ops/eagle.py b/python/sglang/srt/speculative/triton_ops/eagle.py index 4e8df6c60..bd80d9a6e 100644 --- a/python/sglang/srt/speculative/triton_ops/eagle.py +++ b/python/sglang/srt/speculative/triton_ops/eagle.py @@ -7,7 +7,7 @@ def fill_bonus_tokens( accept_tokens, accept_lens, bonus_tokens_ptr, - num_draft_tokens: tl.constexpr, + accept_stride: tl.constexpr, ): # NOTE: we cannot fuse any in-place operations of `accept_lens` inside this kernel # because this kernel reads accept_lens @@ -15,7 +15,8 @@ def fill_bonus_tokens( # `accept_lens` includes the bonus token; the last accepted slot is at -1. accept_len = tl.load(accept_lens + pid) - bonus_token_idx = num_draft_tokens * pid + accept_len - 1 + # accept_stride = per-req width of accept_tokens (= accept_index.shape[1]). + bonus_token_idx = accept_stride * pid + accept_len - 1 bonus_token = tl.load(accept_tokens + bonus_token_idx) tl.store(bonus_tokens_ptr + pid, bonus_token) diff --git a/test/registered/spec/eagle/test_spec_eagle_stress.py b/test/registered/spec/eagle/test_spec_eagle_stress.py index eff30e1db..af4730377 100644 --- a/test/registered/spec/eagle/test_spec_eagle_stress.py +++ b/test/registered/spec/eagle/test_spec_eagle_stress.py @@ -20,7 +20,7 @@ from sglang.test.kits.spec_server_kits import ( ) from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base -register_cuda_ci(est_time=600, stage="base-b", runner_config="1-gpu-large") +register_cuda_ci(est_time=780, stage="base-b", runner_config="1-gpu-large") class TestEagle3Perf(Eagle3Base, SpecPerfKit): @@ -40,6 +40,20 @@ class TestEagleLlama2Retract(EagleLlama2Base, SpecAccuracyKit, SpecFeatureKit): ) +class TestEagle3Topk16V2Retract(Eagle3Base, SpecAccuracyKit, SpecFeatureKit): + """EAGLE3 topk=16 tree on spec v2 under retract; must not leak KV. Stresses + the accepted-path KV move (move_accepted_tokens_to_target_kvcache).""" + + spec_topk = 16 + spec_tokens = 64 + disable_overlap = False + cuda_graph_max_bs = 5 + max_running_requests = 64 + gsm8k_accept_len_thres = 2.4 + extra_args = ("--max-total-tokens", 4500) # small KV to trigger retract + env_overrides = ((envs.SGLANG_TEST_RETRACT, True),) + + class TestEagleLlama2AbortAll(EagleLlama2Base, AbortAllMixin): abort_all_max_new_tokens = 4000 diff --git a/test/registered/spec/eagle/test_spec_eagle_topk.py b/test/registered/spec/eagle/test_spec_eagle_topk.py index 2b3914bba..369a5bb7f 100644 --- a/test/registered/spec/eagle/test_spec_eagle_topk.py +++ b/test/registered/spec/eagle/test_spec_eagle_topk.py @@ -1,7 +1,9 @@ """topk > 1 tree drafting (EAGLE3 topk16 + EAGLE/Llama-2 topk8). -topk > 1 always routes to spec v1; flashinfer is pinned (topk > 1 can't use fa3). -Runs on the cheap (5090) runner -- functional sanity only, no perf/stress. +topk > 1 routes to spec v1, except page_size==1 which can also stay on spec v2 +(overlap). flashinfer is pinned because this runs on the cheap (5090) runner, +where fa3 (Hopper-only) isn't available -- functional sanity only, no perf/stress. +(topk > 1 on fa3 is covered on the Hopper runner in test_spec_eagle_fa3.py.) """ import unittest @@ -16,7 +18,7 @@ from sglang.test.kits.spec_server_kits import ( ) from sglang.test.server_fixtures.spec_eagle_fixture import Eagle3Base, EagleLlama2Base -register_cuda_ci(est_time=840, stage="base-b", runner_config="1-gpu-small") +register_cuda_ci(est_time=1180, stage="base-b", runner_config="1-gpu-small") class TestEagle3Topk16(Eagle3Base, SpecCorrectnessKit, SpecAccuracyKit, SpecLogprobKit): @@ -31,6 +33,13 @@ class TestEagle3Topk16(Eagle3Base, SpecCorrectnessKit, SpecAccuracyKit, SpecLogp gsm8k_accept_len_thres = 2.4 # EAGLE3 topk16 gsm8k accept ~2.48 +class TestEagle3Topk16SpecV2(TestEagle3Topk16, SpecFeatureKit): + """EAGLE3 topk=16 tree on spec v2 (overlap, page1): guards the v2 tree path's + accepted-path compaction, validated by logprob_spec_v2_match.""" + + disable_overlap = False + + class TestEagleLlama2Suite( EagleLlama2Base, SpecCorrectnessKit,