diff --git a/python/sglang/srt/batch_overlap/two_batch_overlap.py b/python/sglang/srt/batch_overlap/two_batch_overlap.py index 4568e3410..14da7ad9e 100644 --- a/python/sglang/srt/batch_overlap/two_batch_overlap.py +++ b/python/sglang/srt/batch_overlap/two_batch_overlap.py @@ -220,24 +220,26 @@ def split_spec_info( positions = spec_info.positions[start_token_index:end_token_index] else: positions = None - if spec_info.retrive_index is not None: - retrive_index = spec_info.retrive_index[start_seq_index:end_seq_index] + if spec_info.retrieve_index is not None: + retrieve_index = spec_info.retrieve_index[start_seq_index:end_seq_index] else: - retrive_index = None - if spec_info.retrive_next_token is not None: - retrive_next_token = spec_info.retrive_next_token[start_seq_index:end_seq_index] - else: - retrive_next_token = None - if spec_info.retrive_next_sibling is not None: - retrive_next_sibling = spec_info.retrive_next_sibling[ + retrieve_index = None + if spec_info.retrieve_next_token is not None: + retrieve_next_token = spec_info.retrieve_next_token[ start_seq_index:end_seq_index ] else: - retrive_next_sibling = None - if spec_info.retrive_cum_len is not None: - retrive_cum_len = spec_info.retrive_cum_len[start_seq_index:end_seq_index] + retrieve_next_token = None + if spec_info.retrieve_next_sibling is not None: + retrieve_next_sibling = spec_info.retrieve_next_sibling[ + start_seq_index:end_seq_index + ] else: - retrive_cum_len = None + retrieve_next_sibling = None + if spec_info.retrieve_cum_len is not None: + retrieve_cum_len = spec_info.retrieve_cum_len[start_seq_index:end_seq_index] + else: + retrieve_cum_len = None if spec_info.seq_lens_cpu is not None: seq_lens_cpu = spec_info.seq_lens_cpu[start_seq_index:end_seq_index] @@ -252,10 +254,10 @@ def split_spec_info( custom_mask=custom_mask, draft_token=draft_token, positions=positions, - retrive_index=retrive_index, - retrive_next_token=retrive_next_token, - retrive_next_sibling=retrive_next_sibling, - retrive_cum_len=retrive_cum_len, + retrieve_index=retrieve_index, + retrieve_next_token=retrieve_next_token, + retrieve_next_sibling=retrieve_next_sibling, + retrieve_cum_len=retrieve_cum_len, seq_lens_cpu=seq_lens_cpu, seq_lens_sum=seq_lens_sum, ) diff --git a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py index aa3c1a817..45a9a4c99 100644 --- a/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py +++ b/python/sglang/srt/layers/attention/hybrid_linear_attn_backend.py @@ -187,8 +187,10 @@ class MambaAttnBackendBase(AttentionBackend): ) if self.topk > 1: - retrieve_next_token = forward_batch.spec_info.retrive_next_token - retrieve_next_sibling = forward_batch.spec_info.retrive_next_sibling + retrieve_next_token = forward_batch.spec_info.retrieve_next_token + retrieve_next_sibling = ( + forward_batch.spec_info.retrieve_next_sibling + ) # retrieve_next_token is None during dummy run so skip tensor creation if retrieve_next_token is not None: retrieve_parent_token = torch.empty_like(retrieve_next_token) @@ -485,8 +487,8 @@ class MambaAttnBackendBase(AttentionBackend): # If topk > 1, we need to use retrieve_next_token and retrieve_next_sibling to handle the eagle tree custom attention mask if forward_mode.is_target_verify() and self.topk > 1: # They are None during cuda graph capture so skip the copy_... - # self.retrieve_next_token_list[bs - 1].copy_(spec_info.retrive_next_token) - # self.retrieve_next_sibling_list[bs - 1].copy_(spec_info.retrive_next_sibling) + # self.retrieve_next_token_list[bs - 1].copy_(spec_info.retrieve_next_token) + # self.retrieve_next_sibling_list[bs - 1].copy_(spec_info.retrieve_next_sibling) return ForwardMetadata( query_start_loc=self.query_start_loc_list[bs - 1], mamba_cache_indices=self.state_indices_list[bs - 1], @@ -545,12 +547,12 @@ class MambaAttnBackendBase(AttentionBackend): # If topk > 1, we need to use retrieve_next_token and retrieve_next_sibling to handle the eagle tree custom attention mask if forward_mode.is_target_verify() and self.topk > 1: - bs_without_pad = spec_info.retrive_next_token.shape[0] + bs_without_pad = spec_info.retrieve_next_token.shape[0] self.retrieve_next_token_list[bs - 1][:bs_without_pad].copy_( - spec_info.retrive_next_token + spec_info.retrieve_next_token ) self.retrieve_next_sibling_list[bs - 1][:bs_without_pad].copy_( - spec_info.retrive_next_sibling + spec_info.retrieve_next_sibling ) return ForwardMetadata( query_start_loc=self.query_start_loc_list[bs - 1], diff --git a/python/sglang/srt/model_executor/cpu_graph_runner.py b/python/sglang/srt/model_executor/cpu_graph_runner.py index d945c4837..6628a516b 100644 --- a/python/sglang/srt/model_executor/cpu_graph_runner.py +++ b/python/sglang/srt/model_executor/cpu_graph_runner.py @@ -842,10 +842,10 @@ class CPUGraphRunner: draft_token=None, custom_mask=self.custom_mask, positions=None, - retrive_index=None, - retrive_next_token=None, - retrive_next_sibling=None, - retrive_cum_len=None, + retrieve_index=None, + retrieve_next_token=None, + retrieve_next_sibling=None, + retrieve_cum_len=None, spec_steps=self.model_runner.server_args.speculative_num_steps, topk=self.model_runner.server_args.speculative_eagle_topk, draft_token_num=self.model_runner.server_args.speculative_num_draft_tokens, diff --git a/python/sglang/srt/model_executor/cuda_graph_runner.py b/python/sglang/srt/model_executor/cuda_graph_runner.py index 2d2795275..c56d5e565 100644 --- a/python/sglang/srt/model_executor/cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/cuda_graph_runner.py @@ -1334,10 +1334,10 @@ class CudaGraphRunner: draft_token=None, custom_mask=self.buffers.custom_mask, positions=None, - retrive_index=None, - retrive_next_token=None, - retrive_next_sibling=None, - retrive_cum_len=None, + retrieve_index=None, + retrieve_next_token=None, + retrieve_next_sibling=None, + retrieve_cum_len=None, spec_steps=self.speculative_num_steps, topk=self.model_runner.server_args.speculative_eagle_topk, draft_token_num=self.speculative_num_draft_tokens, @@ -1379,9 +1379,9 @@ class CudaGraphRunner: draft_token=None, tree_mask=self.buffers.custom_mask, positions=None, - retrive_index=None, - retrive_next_token=None, - retrive_next_sibling=None, + retrieve_index=None, + retrieve_next_token=None, + retrieve_next_sibling=None, draft_token_num=self.num_tokens_per_bs, ) spec_info.capture_hidden_mode = CaptureHiddenMode.NULL diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index e94ca7318..4304bb064 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -2393,10 +2393,10 @@ class ModelRunner(ModelRunnerKVCacheMixin): draft_token=None, custom_mask=buffers.custom_mask, positions=None, - retrive_index=None, - retrive_next_token=None, - retrive_next_sibling=None, - retrive_cum_len=None, + retrieve_index=None, + retrieve_next_token=None, + retrieve_next_sibling=None, + retrieve_cum_len=None, spec_steps=self.server_args.speculative_num_steps, topk=self.server_args.speculative_eagle_topk, draft_token_num=self.server_args.speculative_num_draft_tokens, @@ -2427,9 +2427,9 @@ class ModelRunner(ModelRunnerKVCacheMixin): draft_token=None, tree_mask=buffers.custom_mask, positions=None, - retrive_index=None, - retrive_next_token=None, - retrive_next_sibling=None, + retrieve_index=None, + retrieve_next_token=None, + retrieve_next_sibling=None, draft_token_num=num_tokens_per_bs, ) spec_info.capture_hidden_mode = CaptureHiddenMode.NULL diff --git a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py index 7c2eebf57..31dac386d 100644 --- a/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/piecewise_cuda_graph_runner.py @@ -832,10 +832,10 @@ class PiecewiseCudaGraphRunner: draft_token=None, custom_mask=self.custom_mask, positions=None, - retrive_index=None, - retrive_next_token=None, - retrive_next_sibling=None, - retrive_cum_len=None, + retrieve_index=None, + retrieve_next_token=None, + retrieve_next_sibling=None, + retrieve_cum_len=None, spec_steps=self.model_runner.server_args.speculative_num_steps, topk=self.model_runner.server_args.speculative_eagle_topk, draft_token_num=self.model_runner.server_args.speculative_num_draft_tokens, diff --git a/python/sglang/srt/speculative/dflash_utils.py b/python/sglang/srt/speculative/dflash_utils.py index ddec049e0..4f0d10733 100644 --- a/python/sglang/srt/speculative/dflash_utils.py +++ b/python/sglang/srt/speculative/dflash_utils.py @@ -618,6 +618,7 @@ def compute_dflash_sampling_accept_len_and_bonus( accept_index=accept_index, accept_token_num=accept_token_num, candidates=candidates_i64, + # kwarg LHS retained as `retrive_*` to match sgl_kernel op schema. retrive_index=retrieve_index, retrive_next_token=retrieve_next_token, retrive_next_sibling=retrieve_next_sibling, diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index a21f6b0ae..8cb4bd5ce 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -61,10 +61,10 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): draft_token: torch.Tensor custom_mask: torch.Tensor positions: torch.Tensor - retrive_index: torch.Tensor - retrive_next_token: torch.Tensor - retrive_next_sibling: torch.Tensor - retrive_cum_len: torch.Tensor + retrieve_index: torch.Tensor + retrieve_next_token: torch.Tensor + retrieve_next_sibling: torch.Tensor + retrieve_cum_len: torch.Tensor spec_steps: int topk: int draft_token_num: int @@ -88,16 +88,16 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): draft_token=torch.empty((0,), dtype=torch.long, device="cuda"), custom_mask=torch.full((0,), True, dtype=torch.bool, device="cuda"), positions=torch.empty((0,), dtype=torch.int64, device="cuda"), - retrive_index=torch.full( + retrieve_index=torch.full( (0, num_verify_tokens), -1, dtype=torch.long, device="cuda" ), - retrive_next_token=torch.full( + retrieve_next_token=torch.full( (0, num_verify_tokens), -1, dtype=torch.long, device="cuda" ), - retrive_next_sibling=torch.full( + retrieve_next_sibling=torch.full( (0, num_verify_tokens), -1, dtype=torch.long, device="cuda" ), - retrive_cum_len=None, + retrieve_cum_len=None, topk=topk, draft_token_num=num_verify_tokens, spec_steps=spec_steps, @@ -258,7 +258,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): ), ) - bs = self.retrive_index.shape[0] + bs = self.retrieve_index.shape[0] candidates = self.draft_token.reshape(bs, self.draft_token_num) sampling_info = batch.sampling_info @@ -272,8 +272,10 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): if bs != len(sampling_info): sampling_info = copy.deepcopy(sampling_info) - # NOTE: retrive_index are the indices of the requests that are kept. - sampling_info.filter_batch(self.retrive_index.tolist(), self.retrive_index) + # NOTE: retrieve_index are the indices of the requests that are kept. + sampling_info.filter_batch( + self.retrieve_index.tolist(), self.retrieve_index + ) # Apply the custom logit processors if registered in the sampling info. if sampling_info.has_custom_logit_processor: @@ -322,9 +324,9 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): accept_index=accept_index, # mutable accept_token_num=accept_length, # mutable candidates=candidates, - retrive_index=self.retrive_index, - retrive_next_token=self.retrive_next_token, - retrive_next_sibling=self.retrive_next_sibling, + 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.topk, ) @@ -370,9 +372,10 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin): accept_index=accept_index, # mutable accept_token_num=accept_length, # mutable candidates=candidates, - retrive_index=self.retrive_index, - retrive_next_token=self.retrive_next_token, - retrive_next_sibling=self.retrive_next_sibling, + # 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, diff --git a/python/sglang/srt/speculative/eagle_info_v2.py b/python/sglang/srt/speculative/eagle_info_v2.py index 473883bed..75dbf6a15 100644 --- a/python/sglang/srt/speculative/eagle_info_v2.py +++ b/python/sglang/srt/speculative/eagle_info_v2.py @@ -373,9 +373,9 @@ class EagleVerifyInputV2Mixin: accept_index=accept_index, # mutable accept_token_num=accept_length, # mutable candidates=candidates, - retrive_index=self.retrive_index, - retrive_next_token=self.retrive_next_token, - retrive_next_sibling=self.retrive_next_sibling, + 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.topk, ) @@ -415,9 +415,10 @@ class EagleVerifyInputV2Mixin: accept_index=accept_index, # mutable accept_token_num=accept_length, # mutable candidates=candidates, - retrive_index=self.retrive_index, - retrive_next_token=self.retrive_next_token, - retrive_next_sibling=self.retrive_next_sibling, + # 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, diff --git a/python/sglang/srt/speculative/eagle_utils.py b/python/sglang/srt/speculative/eagle_utils.py index f41a92523..8ce0a27b0 100644 --- a/python/sglang/srt/speculative/eagle_utils.py +++ b/python/sglang/srt/speculative/eagle_utils.py @@ -104,10 +104,10 @@ def build_tree_kernel_efficient( raise NotImplementedError(f"Invalid tree mask: {tree_mask_mode=}") # TODO: make them torch.empty and fuse them into `sgl_build_tree_kernel` - retrive_buf = torch.full( + retrieve_buf = torch.full( (3, bs, num_verify_tokens), -1, device=device, dtype=torch.long ) - retrive_index, retrive_next_token, retrive_next_sibling = retrive_buf + retrieve_index, retrieve_next_token, retrieve_next_sibling = retrieve_buf # position: where each token belongs to # e.g. if depth of each draft token is [0, 1, 1, 2] and the prompt length is 7 # then, positions = [7, 8, 8, 9] @@ -125,9 +125,9 @@ def build_tree_kernel_efficient( seq_lens, tree_mask, positions, - retrive_index, - retrive_next_token, - retrive_next_sibling, + retrieve_index, + retrieve_next_token, + retrieve_next_sibling, topk, spec_steps, num_verify_tokens, @@ -140,9 +140,9 @@ def build_tree_kernel_efficient( seq_lens, tree_mask, positions, - retrive_index, - retrive_next_token, - retrive_next_sibling, + retrieve_index, + retrieve_next_token, + retrieve_next_sibling, topk, spec_steps, num_verify_tokens, @@ -151,9 +151,9 @@ def build_tree_kernel_efficient( return ( tree_mask, positions, - retrive_index, - retrive_next_token, - retrive_next_sibling, + retrieve_index, + retrieve_next_token, + retrieve_next_sibling, draft_tokens, ) @@ -163,9 +163,9 @@ def verify_tree_greedy_func( accept_index: torch.Tensor, accept_token_num: torch.Tensor, candidates: torch.Tensor, - retrive_index: torch.Tensor, - retrive_next_token: torch.Tensor, - retrive_next_sibling: torch.Tensor, + retrieve_index: torch.Tensor, + retrieve_next_token: torch.Tensor, + retrieve_next_sibling: torch.Tensor, target_predict: torch.Tensor, topk: int = -1, ): @@ -177,9 +177,10 @@ def verify_tree_greedy_func( accept_index=accept_index, # mutable accept_token_num=accept_token_num, # mutable candidates=candidates, - retrive_index=retrive_index, - retrive_next_token=retrive_next_token, - retrive_next_sibling=retrive_next_sibling, + # kwarg LHS retained as `retrive_*` to match sgl_kernel op schema. + retrive_index=retrieve_index, + retrive_next_token=retrieve_next_token, + retrive_next_sibling=retrieve_next_sibling, target_predict=target_predict, ) @@ -191,9 +192,10 @@ def verify_tree_greedy_func( accept_index=accept_index, accept_token_num=accept_token_num, candidates=candidates, - retrive_index=retrive_index, - retrive_next_token=retrive_next_token, - retrive_next_sibling=retrive_next_sibling, + # kwarg LHS retained as `retrive_*` to match sgl_kernel op schema. + retrive_index=retrieve_index, + retrive_next_token=retrieve_next_token, + retrive_next_sibling=retrieve_next_sibling, target_predict=target_predict, ) return predicts, accept_index, accept_token_num diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index 80f5266c8..cf99854fb 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -763,9 +763,9 @@ class EAGLEWorker(TpModelWorker): ( tree_mask, position, - retrive_index, - retrive_next_token, - retrive_next_sibling, + retrieve_index, + retrieve_next_token, + retrieve_next_sibling, draft_tokens, ) = build_tree_kernel_efficient( spec_info.verified_id, @@ -783,10 +783,10 @@ class EAGLEWorker(TpModelWorker): draft_token=draft_tokens, custom_mask=tree_mask, positions=position, - retrive_index=retrive_index, - retrive_next_token=retrive_next_token, - retrive_next_sibling=retrive_next_sibling, - retrive_cum_len=None, + retrieve_index=retrieve_index, + retrieve_next_token=retrieve_next_token, + retrieve_next_sibling=retrieve_next_sibling, + retrieve_cum_len=None, spec_steps=self.speculative_num_steps, topk=self.topk, draft_token_num=self.speculative_num_draft_tokens, @@ -897,10 +897,10 @@ class EAGLEWorker(TpModelWorker): assert model_worker_batch.capture_hidden_mode == spec_info.capture_hidden_mode if batch.has_grammar: - retrieve_next_token_cpu = spec_info.retrive_next_token.cpu() - retrieve_next_sibling_cpu = spec_info.retrive_next_sibling.cpu() + retrieve_next_token_cpu = spec_info.retrieve_next_token.cpu() + retrieve_next_sibling_cpu = spec_info.retrieve_next_sibling.cpu() draft_tokens_cpu = spec_info.draft_token.view( - spec_info.retrive_next_token.shape + spec_info.retrieve_next_token.shape ).cpu() # Forward @@ -927,7 +927,7 @@ class EAGLEWorker(TpModelWorker): if vocab_mask is not None: assert spec_info.grammar is not None - vocab_mask = vocab_mask.to(spec_info.retrive_next_token.device) + vocab_mask = vocab_mask.to(spec_info.retrieve_next_token.device) # NOTE (sk): otherwise, this vocab mask will be the one from the previous extend stage # and will be applied to produce wrong results batch.sampling_info.vocab_mask = None diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index e7573b422..01a07ad53 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -364,9 +364,9 @@ class EagleDraftWorker(BaseDraftWorker): ( tree_mask, position, - retrive_index, - retrive_next_token, - retrive_next_sibling, + retrieve_index, + retrieve_next_token, + retrieve_next_sibling, draft_tokens, ) = build_tree_kernel_efficient( draft_input.verified_id, @@ -387,10 +387,10 @@ class EagleDraftWorker(BaseDraftWorker): draft_token=draft_tokens, custom_mask=tree_mask, positions=position, - retrive_index=retrive_index, - retrive_next_token=retrive_next_token, - retrive_next_sibling=retrive_next_sibling, - retrive_cum_len=None, + retrieve_index=retrieve_index, + retrieve_next_token=retrieve_next_token, + retrieve_next_sibling=retrieve_next_sibling, + retrieve_cum_len=None, spec_steps=self.speculative_num_steps, topk=self.topk, draft_token_num=self.speculative_num_draft_tokens, @@ -786,10 +786,10 @@ class EAGLEWorkerV2(BaseSpecWorker): # Prepare grammar data on CPU if needed if batch.has_grammar: - retrieve_next_token_cpu = verify_input.retrive_next_token.cpu() - retrieve_next_sibling_cpu = verify_input.retrive_next_sibling.cpu() + retrieve_next_token_cpu = verify_input.retrieve_next_token.cpu() + retrieve_next_sibling_cpu = verify_input.retrieve_next_sibling.cpu() draft_tokens_cpu = verify_input.draft_token.view( - verify_input.retrive_next_token.shape + verify_input.retrieve_next_token.shape ).cpu() # Run target verify batch in the main compute stream (GPU compute) @@ -816,7 +816,7 @@ class EAGLEWorkerV2(BaseSpecWorker): if vocab_mask is not None: assert verify_input.grammar is not None - vocab_mask = vocab_mask.to(verify_input.retrive_next_token.device) + vocab_mask = vocab_mask.to(verify_input.retrieve_next_token.device) # NOTE: otherwise, this vocab mask will be the one from the previous extend stage # and will be applied to produce wrong results batch.sampling_info.vocab_mask = None diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker.py b/python/sglang/srt/speculative/multi_layer_eagle_worker.py index fb3434970..ec05355e6 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker.py @@ -435,9 +435,9 @@ class MultiLayerEagleWorker(TpModelWorker): ( tree_mask, position, - retrive_index, - retrive_next_token, - retrive_next_sibling, + retrieve_index, + retrieve_next_token, + retrieve_next_sibling, draft_tokens, ) = build_tree_kernel_efficient( spec_info.verified_id, @@ -455,10 +455,10 @@ class MultiLayerEagleWorker(TpModelWorker): draft_token=draft_tokens, custom_mask=tree_mask, positions=position, - retrive_index=retrive_index, - retrive_next_token=retrive_next_token, - retrive_next_sibling=retrive_next_sibling, - retrive_cum_len=None, + retrieve_index=retrieve_index, + retrieve_next_token=retrieve_next_token, + retrieve_next_sibling=retrieve_next_sibling, + retrieve_cum_len=None, spec_steps=self.speculative_num_steps, topk=self.topk, draft_token_num=self.server_args.speculative_num_draft_tokens, @@ -488,10 +488,10 @@ class MultiLayerEagleWorker(TpModelWorker): model_worker_batch.return_hidden_states_before_norm = True if batch.has_grammar: - retrieve_next_token_cpu = spec_info.retrive_next_token.cpu() - retrieve_next_sibling_cpu = spec_info.retrive_next_sibling.cpu() + retrieve_next_token_cpu = spec_info.retrieve_next_token.cpu() + retrieve_next_sibling_cpu = spec_info.retrieve_next_sibling.cpu() draft_tokens_cpu = spec_info.draft_token.view( - spec_info.retrive_next_token.shape + spec_info.retrieve_next_token.shape ).cpu() # Forward @@ -518,7 +518,7 @@ class MultiLayerEagleWorker(TpModelWorker): if vocab_mask is not None: assert spec_info.grammar is not None - vocab_mask = vocab_mask.to(spec_info.retrive_next_token.device) + vocab_mask = vocab_mask.to(spec_info.retrieve_next_token.device) # NOTE (sk): otherwise, this vocab mask will be the one from the previous extend stage # and will be applied to produce wrong results batch.sampling_info.vocab_mask = None 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 89d737dc9..40eb2ad66 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -248,9 +248,9 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker): ( tree_mask, position, - retrive_index, - retrive_next_token, - retrive_next_sibling, + retrieve_index, + retrieve_next_token, + retrieve_next_sibling, draft_tokens, ) = build_tree_kernel_efficient( draft_input.verified_id, @@ -271,10 +271,10 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker): draft_token=draft_tokens, custom_mask=tree_mask, positions=position, - retrive_index=retrive_index, - retrive_next_token=retrive_next_token, - retrive_next_sibling=retrive_next_sibling, - retrive_cum_len=None, + retrieve_index=retrieve_index, + retrieve_next_token=retrieve_next_token, + retrieve_next_sibling=retrieve_next_sibling, + retrieve_cum_len=None, spec_steps=self.speculative_num_steps, topk=self.topk, draft_token_num=self.speculative_num_draft_tokens, diff --git a/python/sglang/srt/speculative/ngram_info.py b/python/sglang/srt/speculative/ngram_info.py index bcfad4bae..e6baf7ed9 100644 --- a/python/sglang/srt/speculative/ngram_info.py +++ b/python/sglang/srt/speculative/ngram_info.py @@ -54,9 +54,9 @@ class NgramVerifyInput(SpecInput): draft_token: torch.Tensor, tree_mask: torch.Tensor, positions: torch.Tensor, - retrive_index: torch.Tensor, - retrive_next_token: torch.Tensor, - retrive_next_sibling: torch.Tensor, + retrieve_index: torch.Tensor, + retrieve_next_token: torch.Tensor, + retrieve_next_sibling: torch.Tensor, draft_token_num: int, grammar: BaseGrammarObject = None, ): @@ -64,9 +64,9 @@ class NgramVerifyInput(SpecInput): self.draft_token = draft_token self.custom_mask = tree_mask self.positions = positions - self.retrive_index = retrive_index - self.retrive_next_token = retrive_next_token - self.retrive_next_sibling = retrive_next_sibling + self.retrieve_index = retrieve_index + self.retrieve_next_token = retrieve_next_token + self.retrieve_next_sibling = retrieve_next_sibling self.draft_token_num = draft_token_num self.device = self.custom_mask.device self.grammar = grammar @@ -303,9 +303,9 @@ class NgramVerifyInput(SpecInput): accept_index=self.accepted_indices, # mutable accept_token_num=self.accept_length, # mutable candidates=candidates, - retrive_index=self.retrive_index, - retrive_next_token=self.retrive_next_token, - retrive_next_sibling=self.retrive_next_sibling, + retrieve_index=self.retrieve_index, + retrieve_next_token=self.retrieve_next_token, + retrieve_next_sibling=self.retrieve_next_sibling, target_predict=target_predict, ) @@ -365,9 +365,10 @@ class NgramVerifyInput(SpecInput): accept_index=self.accepted_indices, # mutable accept_token_num=self.accept_length, # mutable candidates=candidates.to(torch.int64), - retrive_index=self.retrive_index.to(torch.int64), - retrive_next_token=self.retrive_next_token.to(torch.int64), - retrive_next_sibling=self.retrive_next_sibling.to(torch.int64), + # kwarg LHS retained as `retrive_*` to match sgl_kernel op schema. + retrive_index=self.retrieve_index.to(torch.int64), + retrive_next_token=self.retrieve_next_token.to(torch.int64), + retrive_next_sibling=self.retrieve_next_sibling.to(torch.int64), uniform_samples=coins, uniform_samples_for_final_sampling=coins_for_final_sampling, target_probs=target_probs, @@ -384,13 +385,15 @@ class NgramVerifyInput(SpecInput): page_size: int, vocab_mask: Optional[torch.Tensor] = None, # For grammar ) -> torch.Tensor: - bs = self.retrive_index.shape[0] + bs = self.retrieve_index.shape[0] sampling_info = batch.sampling_info if bs != len(sampling_info): sampling_info = copy.deepcopy(sampling_info) - # NOTE: retrive_index are the indices of the requests that are kept. - sampling_info.filter_batch(self.retrive_index.tolist(), self.retrive_index) + # NOTE: retrieve_index are the indices of the requests that are kept. + sampling_info.filter_batch( + self.retrieve_index.tolist(), self.retrieve_index + ) # Apply the custom logit processors if registered in the sampling info. if sampling_info.has_custom_logit_processor: diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index 6dec89584..47724ef28 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -117,12 +117,12 @@ class NGRAMWorker: dtype=torch.int64, device=self.device, ) - self.retrive_next_token = torch.empty( + self.retrieve_next_token = torch.empty( (self.max_batch_size, self.draft_token_num), dtype=torch.int64, device=self.device, ) - self.retrive_next_sibling = torch.empty( + self.retrieve_next_sibling = torch.empty( (self.max_batch_size, self.draft_token_num), dtype=torch.int64, device=self.device, @@ -137,14 +137,14 @@ class NGRAMWorker: self.draft_tokens_batch = [] self.tree_mask_batch = [] self.retrieve_indexes_batch = [] - self.retrive_next_token_batch = [] - self.retrive_next_sibling_batch = [] + self.retrieve_next_token_batch = [] + self.retrieve_next_sibling_batch = [] self.positions_batch = [] for bs in range(0, self.max_batch_size + 1): self.retrieve_indexes_batch.append(self.retrieve_indexes[:bs, :]) - self.retrive_next_token_batch.append(self.retrive_next_token[:bs, :]) - self.retrive_next_sibling_batch.append(self.retrive_next_sibling[:bs, :]) + self.retrieve_next_token_batch.append(self.retrieve_next_token[:bs, :]) + self.retrieve_next_sibling_batch.append(self.retrieve_next_sibling[:bs, :]) self.positions_batch.append(self.positions[: bs * self.draft_token_num]) self.draft_tokens_batch.append( self.draft_tokens[: bs * self.draft_token_num] @@ -186,9 +186,9 @@ class NGRAMWorker: bs = batch.batch_size() - retrive_index = self.retrieve_indexes_batch[bs] - retrive_next_token = self.retrive_next_token_batch[bs] - retrive_next_sibling = self.retrive_next_sibling_batch[bs] + retrieve_index = self.retrieve_indexes_batch[bs] + retrieve_next_token = self.retrieve_next_token_batch[bs] + retrieve_next_sibling = self.retrieve_next_sibling_batch[bs] positions = self.positions_batch[bs] tree_mask = self.tree_mask_batch[bs] draft_tokens = self.draft_tokens_batch[bs] @@ -201,9 +201,9 @@ class NGRAMWorker: tree_mask, batch.seq_lens, positions, # mutable - retrive_index, # mutable - retrive_next_token, # mutable - retrive_next_sibling, # mutable + retrieve_index, # mutable + retrieve_next_token, # mutable + retrieve_next_sibling, # mutable bs, self.draft_token_num, ) @@ -230,9 +230,9 @@ class NGRAMWorker: draft_tokens, tree_mask, positions, - retrive_index, - retrive_next_token, - retrive_next_sibling, + retrieve_index, + retrieve_next_token, + retrieve_next_sibling, self.draft_token_num, ) batch.spec_info.prepare_for_verify(batch, self.page_size) @@ -266,10 +266,10 @@ class NGRAMWorker: if model_worker_batch.forward_mode.is_target_verify(): if batch.has_grammar: - retrieve_next_token_cpu = spec_info.retrive_next_token.cpu() - retrieve_next_sibling_cpu = spec_info.retrive_next_sibling.cpu() + retrieve_next_token_cpu = spec_info.retrieve_next_token.cpu() + retrieve_next_sibling_cpu = spec_info.retrieve_next_sibling.cpu() draft_tokens_cpu = spec_info.draft_token.view( - spec_info.retrive_next_token.shape + spec_info.retrieve_next_token.shape ).cpu() set_time_batch(batch.reqs, "set_spec_verify_start_time", trace_only=True) @@ -298,7 +298,7 @@ class NGRAMWorker: if vocab_mask is not None: assert verify_input.grammar is not None - vocab_mask = vocab_mask.to(verify_input.retrive_next_token.device) + vocab_mask = vocab_mask.to(verify_input.retrieve_next_token.device) # NOTE (sk): otherwise, this vocab mask will be the one from the previous extend stage # and will be applied to produce wrong results batch.sampling_info.vocab_mask = None