diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index 2e4056dfa..73de4be2f 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -8,10 +8,7 @@ from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject from sglang.srt.environ import envs from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode -from sglang.srt.speculative.eagle_info_v2 import ( - EagleDraftInputV2Mixin, - EagleVerifyInputV2Mixin, -) +from sglang.srt.speculative.eagle_info_v2 import EagleDraftInputV2Mixin from sglang.srt.speculative.spec_info import SpecInput, SpecInputType logger = logging.getLogger(__name__) @@ -30,7 +27,7 @@ def _draft_runner_of(worker): @dataclass -class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): +class EagleVerifyInput(SpecInput): draft_token: torch.Tensor custom_mask: torch.Tensor positions: torch.Tensor @@ -54,6 +51,19 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): if self.num_tokens_per_req < 0: self.num_tokens_per_req = self.draft_token_num + @property + def max_tree_depth(self) -> int: + """Longest root-to-leaf chain of the verify tree, incl. the root; + bounds the accept_index row width. EAGLE trees are depth-bounded by + the draft loop. Algorithms with other tree shapes override this.""" + return self.spec_steps + 1 + + @property + def tree_topk(self) -> int: + """Branching factor passed to the tree-verify kernels; -1 means an + irregular tree (no fixed per-level branching).""" + return self.topk + def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]: return self.draft_token_num, self.draft_token_num diff --git a/python/sglang/srt/speculative/eagle_info_v2.py b/python/sglang/srt/speculative/eagle_info_v2.py index 9ca21a3eb..c75589356 100644 --- a/python/sglang/srt/speculative/eagle_info_v2.py +++ b/python/sglang/srt/speculative/eagle_info_v2.py @@ -4,14 +4,7 @@ from dataclasses import dataclass from typing import TYPE_CHECKING import torch -import torch.nn.functional as F -from sglang.srt.distributed import get_tp_group -from sglang.srt.layers.dp_attention import ( - get_attention_tp_group, - is_dp_attention_enabled, -) -from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.mem_cache.common import ( alloc_paged_token_slots_extend, @@ -19,50 +12,16 @@ from sglang.srt.mem_cache.common import ( get_alloc_reserve_per_decode, get_last_loc, ) -from sglang.srt.mem_cache.memory_pool import ReqToTokenPool -from sglang.srt.model_executor.forward_batch_info import ( - CaptureHiddenMode, - ForwardBatch, - ForwardMode, -) -from sglang.srt.sampling.penaltylib.repetition_penalty import apply_scaling_penalties -from sglang.srt.server_args import get_global_server_args -from sglang.srt.speculative.eagle_utils import verify_tree_greedy_func -from sglang.srt.speculative.spec_utils import ( - SIMULATE_ACC_LEN, - generate_simulated_accept_index, - prepare_mamba_track_for_verify, -) from sglang.srt.speculative.triton_ops.cache_locs import ( assign_extend_cache_locs_func as assign_extend_cache_locs_func, ) from sglang.srt.speculative.triton_ops.eagle import ( fill_bonus_tokens as fill_bonus_tokens, ) -from sglang.srt.utils.async_probe import ( - maybe_detect_nan, - maybe_detect_oob, - sanitize_nan_logits, -) -from sglang.srt.utils.common import is_cuda, is_hip, is_musa, is_npu - -_is_cuda = is_cuda() -_is_hip = is_hip() -_is_npu = is_npu() -_is_musa = is_musa() if TYPE_CHECKING: - from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.speculative.eagle_info import ( EagleDraftInput, - EagleVerifyInput, - ) - -if is_cuda() or is_musa(): - from sgl_kernel import ( - top_k_renorm_prob, - top_p_renorm_prob, - tree_speculative_sampling_target_only, ) @@ -164,247 +123,3 @@ class EagleDraftInputV2Mixin: out_cache_loc, bs, ) - - -@dataclass -class EagleVerifyInputV2Mixin: - @property - def max_tree_depth(self: EagleVerifyInput) -> int: - """Longest root-to-leaf chain of the verify tree, incl. the root; - bounds the accept_index row width. EAGLE trees are depth-bounded by - the draft loop. Algorithms with other tree shapes override this.""" - return self.spec_steps + 1 - - @property - def tree_topk(self: EagleVerifyInput) -> int: - """Branching factor passed to the tree-verify kernels; -1 means an - irregular tree (no fixed per-level branching).""" - return self.topk - - def prepare_for_verify( - self: EagleVerifyInput, - req_to_token_pool: ReqToTokenPool, - batch: ScheduleBatch, - target_worker: TpModelWorker, - ): - if not batch.forward_mode.is_idle(): - # Assign cache locations - bs = len(batch.req_pool_indices) - batch.input_ids = self.draft_token - maybe_detect_oob( - batch.input_ids, - 0, - batch.model_config.vocab_size, - "v2 prepare_for_verify input_ids", - ) - device = batch.device - batch.out_cache_loc = assign_extend_cache_locs_func( - req_pool_indices=batch.req_pool_indices, - req_to_token=req_to_token_pool.req_to_token, - start_offset=batch.seq_lens, - end_offset=batch.seq_lens + self.draft_token_num, - batch_size=bs, - draft_token_num=self.draft_token_num, - device=device, - ) - - prepare_mamba_track_for_verify(batch) - - # TBO's split_spec_info reads these; no-verify-sync leaves both None. - self.seq_lens_cpu = batch.seq_lens_cpu - self.seq_lens_sum = ( - int(batch.seq_lens_cpu.sum()) - if batch.seq_lens_cpu is not None - else None - ) - - # Get a forward batch - batch.forward_mode = ( - ForwardMode.IDLE - if batch.forward_mode.is_idle() - else ForwardMode.TARGET_VERIFY - ) - capture_mode = ( - CaptureHiddenMode.NULL - if target_worker.model_runner.spec_algorithm.is_standalone() - else CaptureHiddenMode.FULL - ) - batch.capture_hidden_mode = capture_mode - verify_forward_batch = ForwardBatch.init_new(batch, target_worker.model_runner) - - # Run attention backend plan and cuda graph preparation - can_run_cuda_graph = bool( - target_worker.model_runner.decode_cuda_graph_runner - and target_worker.model_runner.decode_cuda_graph_runner.can_run( - verify_forward_batch - ) - ) - if can_run_cuda_graph: - target_worker.model_runner.decode_cuda_graph_runner.replay_prepare( - verify_forward_batch - ) - verify_forward_batch.mark_forward_metadata_ready() - # Non-cuda-graph: defer init to forward_extend, which runs after - # `_forward_raw -> prepare_mlp_sync_batch` pads the batch. Initing - # here would use pre-pad shapes and trip DSv4 indexer shape match. - - return verify_forward_batch, can_run_cuda_graph - - def sample( - self: EagleVerifyInput, - batch: ScheduleBatch, - logits_output: LogitsProcessorOutput, - vocab_mask: torch.Tensor = None, - ): - """ - Verify and find accepted tokens based on logits output and batch - (which contains spec decoding information). - """ - device = batch.device - if batch.forward_mode.is_idle(): - predict = torch.empty(0, dtype=torch.int32, device=device) - num_correct_drafts = torch.empty(0, dtype=torch.int32, device=device) - accept_index = torch.empty(0, dtype=torch.int32, device=device) - return predict, num_correct_drafts, accept_index - - bs = len(batch.seq_lens) - sampling_info = batch.sampling_info - next_token_logits = logits_output.next_token_logits - - sanitize_nan_logits(next_token_logits, "verify: target model logits") - - # Apply penalty - # This is a relaxed version of penalties for speculative decoding. - if sampling_info.acc_additive_penalties is not None: - next_token_logits.add_( - torch.repeat_interleave( - sampling_info.acc_additive_penalties, self.draft_token_num, dim=0 - ) - ) - if sampling_info.acc_scaling_penalties is not None: - apply_scaling_penalties( - next_token_logits, - torch.repeat_interleave( - sampling_info.acc_scaling_penalties, self.draft_token_num, dim=0 - ), - ) - if sampling_info.logit_bias is not None: - next_token_logits.add_( - torch.repeat_interleave( - sampling_info.logit_bias, self.draft_token_num, dim=0 - ) - ) - - # Apply grammar mask if provided - if vocab_mask is not None: - assert self.grammar is not None - self.grammar.apply_vocab_mask( - logits=next_token_logits, vocab_mask=vocab_mask - ) - - candidates = self.draft_token.reshape(bs, self.draft_token_num) - predict_shape = list(next_token_logits.shape)[:-1] - predict = torch.zeros(predict_shape, dtype=torch.int32, device=device).flatten() - accept_index = torch.full( - (bs, self.max_tree_depth), -1, dtype=torch.int32, device=device - ) - num_correct_drafts = torch.empty((bs,), dtype=torch.int32, device=device) - - # Sample tokens - if sampling_info.is_all_greedy or _is_npu or _is_hip: - target_predict = torch.argmax(next_token_logits, dim=-1) - target_predict = target_predict.reshape(bs, self.draft_token_num) - predict, accept_index, num_correct_drafts = verify_tree_greedy_func( - predicts=predict, # mutable - accept_index=accept_index, # mutable - accept_token_num=num_correct_drafts, # mutable - candidates=candidates, - retrieve_index=self.retrieve_index, - retrieve_next_token=self.retrieve_next_token, - retrieve_next_sibling=self.retrieve_next_sibling, - target_predict=target_predict, - topk=self.tree_topk, - ) - else: - # Apply temperature and get target probs - expanded_temperature = torch.repeat_interleave( - sampling_info.temperatures, self.draft_token_num, dim=0 - ) # (bs * num_draft_tokens, 1) - - target_probs = F.softmax( - next_token_logits / expanded_temperature, dim=-1 - ) # (bs * num_draft_tokens, vocab_size) - maybe_detect_nan(target_probs, "v2 verify: target_probs after softmax") - target_probs = top_k_renorm_prob( - target_probs, - torch.repeat_interleave( - sampling_info.top_ks, self.draft_token_num, dim=0 - ), - ) # (bs * num_draft_tokens, vocab_size) - maybe_detect_nan(target_probs, "v2 verify: target_probs after top_k_renorm") - target_probs = top_p_renorm_prob( - target_probs, - torch.repeat_interleave( - sampling_info.top_ps, self.draft_token_num, dim=0 - ), - ) - maybe_detect_nan(target_probs, "v2 verify: target_probs after top_p_renorm") - target_probs = target_probs.reshape(bs, self.draft_token_num, -1) - draft_probs = torch.zeros_like(target_probs) - - # coins for rejection sampling - coins = torch.rand_like(candidates, dtype=torch.float32, device=device) - # coins for final sampling - coins_for_final_sampling = torch.rand( - (bs,), dtype=torch.float32, device=device - ) - - tree_speculative_sampling_target_only( - predicts=predict, # mutable - accept_index=accept_index, # mutable - accept_token_num=num_correct_drafts, # mutable - candidates=candidates, - # kwarg LHS retained as `retrive_*` to match sgl_kernel op schema. - retrive_index=self.retrieve_index, - retrive_next_token=self.retrieve_next_token, - retrive_next_sibling=self.retrieve_next_sibling, - uniform_samples=coins, - uniform_samples_for_final_sampling=coins_for_final_sampling, - target_probs=target_probs, - draft_probs=draft_probs, - threshold_single=get_global_server_args().speculative_accept_threshold_single, - threshold_acc=get_global_server_args().speculative_accept_threshold_acc, - deterministic=True, - ) - - # Sync sampling results across TP ranks: different GPUs may - # produce slightly different target_probs due to floating-point - # non-determinism in softmax/top_k/top_p, causing different - # sampled tokens. Broadcast from rank 0 to ensure consistency. - tp_group = ( - get_attention_tp_group() - if is_dp_attention_enabled() - else get_tp_group() - ) - if tp_group.world_size > 1: - tp_group.broadcast(predict, src=0) - tp_group.broadcast(accept_index, src=0) - tp_group.broadcast(num_correct_drafts, src=0) - - if SIMULATE_ACC_LEN > 0: - # Do simulation. The helper builds (and returns) a replacement - # accept_index of width spec_steps + 1, so pass max_tree_depth - 1 - # to keep the simulated width identical to the real one. - accept_index = generate_simulated_accept_index( - accept_index=accept_index, - predict=predict, # mutable - num_correct_drafts=num_correct_drafts, # mutable - simulate_acc_len=SIMULATE_ACC_LEN, - bs=bs, - spec_steps=self.max_tree_depth - 1, - ) - - # `num_correct_drafts` stays drafts-only inside this function; the returned - # tensor includes the trailing/bonus token via out-of-place +1 so the - # name no longer flips semantics mid-function (naming doc C2). - return predict, num_correct_drafts + 1, accept_index diff --git a/python/sglang/srt/speculative/eagle_utils.py b/python/sglang/srt/speculative/eagle_utils.py index 3bb3cec3a..31330bf30 100644 --- a/python/sglang/srt/speculative/eagle_utils.py +++ b/python/sglang/srt/speculative/eagle_utils.py @@ -10,8 +10,12 @@ from sglang.srt.utils import is_cuda, is_hip, is_musa, is_npu from sglang.srt.utils.async_probe import maybe_detect_oob if TYPE_CHECKING: + from sglang.srt.layers.logits_processor import LogitsProcessorOutput from sglang.srt.managers.schedule_batch import ScheduleBatch + from sglang.srt.managers.tp_worker import TpModelWorker + from sglang.srt.mem_cache.memory_pool import ReqToTokenPool from sglang.srt.model_executor.model_runner import ModelRunner + from sglang.srt.speculative.eagle_info import EagleVerifyInput _is_cuda = is_cuda() _is_hip = is_hip() @@ -272,3 +276,260 @@ def get_draft_hidden_dim(model_runner: ModelRunner) -> int: num_aux = max(len(layer_ids), 1) return base * num_aux return model_runner.model_config.spec_hidden_size + + +def eagle_prepare_for_verify( + verify_input: EagleVerifyInput, + req_to_token_pool: ReqToTokenPool, + batch: ScheduleBatch, + target_worker: TpModelWorker, +): + from sglang.srt.model_executor.forward_batch_info import ( + CaptureHiddenMode, + ForwardBatch, + ForwardMode, + ) + from sglang.srt.speculative.spec_utils import prepare_mamba_track_for_verify + from sglang.srt.speculative.triton_ops.cache_locs import ( + assign_extend_cache_locs_func, + ) + + if not batch.forward_mode.is_idle(): + # Assign cache locations + bs = len(batch.req_pool_indices) + batch.input_ids = verify_input.draft_token + maybe_detect_oob( + batch.input_ids, + 0, + batch.model_config.vocab_size, + "v2 prepare_for_verify input_ids", + ) + device = batch.device + batch.out_cache_loc = assign_extend_cache_locs_func( + req_pool_indices=batch.req_pool_indices, + req_to_token=req_to_token_pool.req_to_token, + start_offset=batch.seq_lens, + end_offset=batch.seq_lens + verify_input.draft_token_num, + batch_size=bs, + draft_token_num=verify_input.draft_token_num, + device=device, + ) + + prepare_mamba_track_for_verify(batch) + + # TBO's split_spec_info reads these; no-verify-sync leaves both None. + verify_input.seq_lens_cpu = batch.seq_lens_cpu + verify_input.seq_lens_sum = ( + int(batch.seq_lens_cpu.sum()) if batch.seq_lens_cpu is not None else None + ) + + # Get a forward batch + batch.forward_mode = ( + ForwardMode.IDLE if batch.forward_mode.is_idle() else ForwardMode.TARGET_VERIFY + ) + capture_mode = ( + CaptureHiddenMode.NULL + if target_worker.model_runner.spec_algorithm.is_standalone() + else CaptureHiddenMode.FULL + ) + batch.capture_hidden_mode = capture_mode + verify_forward_batch = ForwardBatch.init_new(batch, target_worker.model_runner) + + # Run attention backend plan and cuda graph preparation + can_run_cuda_graph = bool( + target_worker.model_runner.decode_cuda_graph_runner + and target_worker.model_runner.decode_cuda_graph_runner.can_run( + verify_forward_batch + ) + ) + if can_run_cuda_graph: + target_worker.model_runner.decode_cuda_graph_runner.replay_prepare( + verify_forward_batch + ) + verify_forward_batch.mark_forward_metadata_ready() + # Non-cuda-graph: defer init to forward_extend, which runs after + # `_forward_raw -> prepare_mlp_sync_batch` pads the batch. Initing + # here would use pre-pad shapes and trip DSv4 indexer shape match. + + return verify_forward_batch, can_run_cuda_graph + + +def eagle_sample( + verify_input: EagleVerifyInput, + batch: ScheduleBatch, + logits_output: LogitsProcessorOutput, + vocab_mask: torch.Tensor = None, +): + """ + Verify and find accepted tokens based on logits output and batch + (which contains spec decoding information). + """ + import torch.nn.functional as F + + from sglang.srt.distributed import get_tp_group + from sglang.srt.layers.dp_attention import ( + get_attention_tp_group, + is_dp_attention_enabled, + ) + from sglang.srt.sampling.penaltylib.repetition_penalty import ( + apply_scaling_penalties, + ) + from sglang.srt.server_args import get_global_server_args + from sglang.srt.speculative.spec_utils import ( + SIMULATE_ACC_LEN, + generate_simulated_accept_index, + ) + from sglang.srt.utils.async_probe import maybe_detect_nan, sanitize_nan_logits + + device = batch.device + if batch.forward_mode.is_idle(): + predict = torch.empty(0, dtype=torch.int32, device=device) + num_correct_drafts = torch.empty(0, dtype=torch.int32, device=device) + accept_index = torch.empty(0, dtype=torch.int32, device=device) + return predict, num_correct_drafts, accept_index + + bs = len(batch.seq_lens) + sampling_info = batch.sampling_info + next_token_logits = logits_output.next_token_logits + + sanitize_nan_logits(next_token_logits, "verify: target model logits") + + # Apply penalty + # This is a relaxed version of penalties for speculative decoding. + if sampling_info.acc_additive_penalties is not None: + next_token_logits.add_( + torch.repeat_interleave( + sampling_info.acc_additive_penalties, + verify_input.draft_token_num, + dim=0, + ) + ) + if sampling_info.acc_scaling_penalties is not None: + apply_scaling_penalties( + next_token_logits, + torch.repeat_interleave( + sampling_info.acc_scaling_penalties, verify_input.draft_token_num, dim=0 + ), + ) + if sampling_info.logit_bias is not None: + next_token_logits.add_( + torch.repeat_interleave( + sampling_info.logit_bias, verify_input.draft_token_num, dim=0 + ) + ) + + # Apply grammar mask if provided + if vocab_mask is not None: + assert verify_input.grammar is not None + verify_input.grammar.apply_vocab_mask( + logits=next_token_logits, vocab_mask=vocab_mask + ) + + candidates = verify_input.draft_token.reshape(bs, verify_input.draft_token_num) + predict_shape = list(next_token_logits.shape)[:-1] + predict = torch.zeros(predict_shape, dtype=torch.int32, device=device).flatten() + accept_index = torch.full( + (bs, verify_input.max_tree_depth), -1, dtype=torch.int32, device=device + ) + num_correct_drafts = torch.empty((bs,), dtype=torch.int32, device=device) + + # Sample tokens + if sampling_info.is_all_greedy or _is_npu or _is_hip: + target_predict = torch.argmax(next_token_logits, dim=-1) + target_predict = target_predict.reshape(bs, verify_input.draft_token_num) + predict, accept_index, num_correct_drafts = verify_tree_greedy_func( + predicts=predict, # mutable + accept_index=accept_index, # mutable + accept_token_num=num_correct_drafts, # mutable + candidates=candidates, + retrieve_index=verify_input.retrieve_index, + retrieve_next_token=verify_input.retrieve_next_token, + retrieve_next_sibling=verify_input.retrieve_next_sibling, + target_predict=target_predict, + topk=verify_input.tree_topk, + ) + else: + from sgl_kernel import ( + top_k_renorm_prob, + top_p_renorm_prob, + tree_speculative_sampling_target_only, + ) + + # Apply temperature and get target probs + expanded_temperature = torch.repeat_interleave( + sampling_info.temperatures, verify_input.draft_token_num, dim=0 + ) # (bs * num_draft_tokens, 1) + + target_probs = F.softmax( + next_token_logits / expanded_temperature, dim=-1 + ) # (bs * num_draft_tokens, vocab_size) + maybe_detect_nan(target_probs, "v2 verify: target_probs after softmax") + target_probs = top_k_renorm_prob( + target_probs, + torch.repeat_interleave( + sampling_info.top_ks, verify_input.draft_token_num, dim=0 + ), + ) # (bs * num_draft_tokens, vocab_size) + maybe_detect_nan(target_probs, "v2 verify: target_probs after top_k_renorm") + target_probs = top_p_renorm_prob( + target_probs, + torch.repeat_interleave( + sampling_info.top_ps, verify_input.draft_token_num, dim=0 + ), + ) + maybe_detect_nan(target_probs, "v2 verify: target_probs after top_p_renorm") + target_probs = target_probs.reshape(bs, verify_input.draft_token_num, -1) + draft_probs = torch.zeros_like(target_probs) + + # coins for rejection sampling + coins = torch.rand_like(candidates, dtype=torch.float32, device=device) + # coins for final sampling + coins_for_final_sampling = torch.rand((bs,), dtype=torch.float32, device=device) + + tree_speculative_sampling_target_only( + predicts=predict, # mutable + accept_index=accept_index, # mutable + accept_token_num=num_correct_drafts, # mutable + candidates=candidates, + # kwarg LHS retained as `retrive_*` to match sgl_kernel op schema. + retrive_index=verify_input.retrieve_index, + retrive_next_token=verify_input.retrieve_next_token, + retrive_next_sibling=verify_input.retrieve_next_sibling, + uniform_samples=coins, + uniform_samples_for_final_sampling=coins_for_final_sampling, + target_probs=target_probs, + draft_probs=draft_probs, + threshold_single=get_global_server_args().speculative_accept_threshold_single, + threshold_acc=get_global_server_args().speculative_accept_threshold_acc, + deterministic=True, + ) + + # Sync sampling results across TP ranks: different GPUs may + # produce slightly different target_probs due to floating-point + # non-determinism in softmax/top_k/top_p, causing different + # sampled tokens. Broadcast from rank 0 to ensure consistency. + tp_group = ( + get_attention_tp_group() if is_dp_attention_enabled() else get_tp_group() + ) + if tp_group.world_size > 1: + tp_group.broadcast(predict, src=0) + tp_group.broadcast(accept_index, src=0) + tp_group.broadcast(num_correct_drafts, src=0) + + if SIMULATE_ACC_LEN > 0: + # Do simulation. The helper builds (and returns) a replacement + # accept_index of width spec_steps + 1, so pass max_tree_depth - 1 + # to keep the simulated width identical to the real one. + accept_index = generate_simulated_accept_index( + accept_index=accept_index, + predict=predict, # mutable + num_correct_drafts=num_correct_drafts, # mutable + simulate_acc_len=SIMULATE_ACC_LEN, + bs=bs, + spec_steps=verify_input.max_tree_depth - 1, + ) + + # `num_correct_drafts` stays drafts-only inside this function; the returned + # tensor includes the trailing/bonus token via out-of-place +1 so the + # name no longer flips semantics mid-function (naming doc C2). + return predict, num_correct_drafts + 1, accept_index diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 75903eacd..299ad3bc4 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -65,6 +65,8 @@ from sglang.srt.speculative.eagle_utils import ( TreeMaskMode, _eagle_prefill_tail_tokens, build_tree_kernel_efficient, + eagle_prepare_for_verify, + eagle_sample, organize_draft_results, per_step_draft_out_cache_loc, ) @@ -1215,7 +1217,8 @@ class EAGLEWorkerV2(BaseSpecWorker): # Batch 1: Target verify # Prepare for target verify in a separate stream with self.plan_stream_ctx: - verify_forward_batch, can_run_cuda_graph = verify_input.prepare_for_verify( + verify_forward_batch, can_run_cuda_graph = eagle_prepare_for_verify( + verify_input, self.req_to_token_pool, batch, self.target_worker, @@ -1264,7 +1267,7 @@ class EAGLEWorkerV2(BaseSpecWorker): # Run target verify batch in the main compute stream (GPU compute). # Metadata init is skipped iff cuda-graph already ran replay_prepare — - # prepare_for_verify marked the batch in exactly that case; the + # eagle_prepare_for_verify marked the batch in exactly that case; the # non-cuda-graph path stays unmarked and gets forward_extend's init # (post-pad). forward_batch_output = self.target_worker.forward_batch_generation( @@ -1301,7 +1304,7 @@ class EAGLEWorkerV2(BaseSpecWorker): predict, accept_lens, accept_index, - ) = verify_input.sample(batch, logits_output, vocab_mask) + ) = eagle_sample(verify_input, batch, logits_output, vocab_mask) new_seq_lens = batch.seq_lens + accept_lens # Update mamba state for hybrid GDN models after verification 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 7b1b17a35..92d99231a 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -49,7 +49,12 @@ from sglang.srt.speculative.eagle_info import ( EagleVerifyInput, ) from sglang.srt.speculative.eagle_info_v2 import fill_bonus_tokens -from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient +from sglang.srt.speculative.eagle_utils import ( + TreeMaskMode, + build_tree_kernel_efficient, + eagle_prepare_for_verify, + eagle_sample, +) from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import ( MultiLayerEagleMultiStepDraftExtendCudaGraphRunner, ) @@ -807,7 +812,8 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): # Batch 1: Target verify # Prepare for target verify in a separate stream with self.plan_stream_ctx: - verify_forward_batch, can_run_cuda_graph = verify_input.prepare_for_verify( + verify_forward_batch, can_run_cuda_graph = eagle_prepare_for_verify( + verify_input, self.req_to_token_pool, batch, self.target_worker, @@ -834,7 +840,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): ), ) # NOTE: metadata init is skipped here unconditionally, although - # prepare_for_verify only plans when cuda-graph replay_prepare ran. + # eagle_prepare_for_verify only plans when cuda-graph replay_prepare ran. # eagle_worker_v2 re-inits the non-graph path instead (post-pad); this # worker has not adopted that fix, so preserve its behavior verbatim. # On NPU with --disable-cuda-graph, non-graph verify needs metadata init @@ -856,7 +862,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): predict, accept_lens, accept_index, - ) = verify_input.sample(batch, logits_output) + ) = eagle_sample(verify_input, batch, logits_output) new_seq_lens = batch.seq_lens + accept_lens if not batch.forward_mode.is_idle(): diff --git a/python/sglang/srt/speculative/ngram_info.py b/python/sglang/srt/speculative/ngram_info.py index fc74b81de..21136b66a 100644 --- a/python/sglang/srt/speculative/ngram_info.py +++ b/python/sglang/srt/speculative/ngram_info.py @@ -6,14 +6,11 @@ import torch from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton -from sglang.srt.speculative.eagle_info_v2 import ( - EagleDraftInputV2Mixin, - EagleVerifyInputV2Mixin, -) +from sglang.srt.speculative.eagle_info_v2 import EagleDraftInputV2Mixin from sglang.srt.speculative.spec_info import SpecInput, SpecInputType -class NgramVerifyInput(SpecInput, EagleDraftInputV2Mixin, EagleVerifyInputV2Mixin): +class NgramVerifyInput(SpecInput, EagleDraftInputV2Mixin): def __init__( self, draft_token: torch.Tensor = None, diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index b762641ad..4cbcb7f0e 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -14,6 +14,7 @@ from sglang.srt.observability.req_time_stats import set_time_batch from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase from sglang.srt.speculative.cpp_ngram.ngram_corpus import NgramCorpus +from sglang.srt.speculative.eagle_utils import eagle_sample from sglang.srt.speculative.ngram_info import NgramVerifyInput from sglang.srt.speculative.spec_utils import ( commit_mamba_states_after_verify, @@ -432,7 +433,7 @@ class NGRAMWorker(BaseSpecWorker): predict, accept_lens, accept_index, - ) = verify_input.sample(batch, logits_output, vocab_mask) + ) = eagle_sample(verify_input, batch, logits_output, vocab_mask) new_seq_lens = batch.seq_lens + accept_lens commit_mamba_states_after_verify( self.target_worker,