diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index a168481c5..53cbbd1fd 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -25,7 +25,6 @@ from sglang.kernels.ops.activation.softcap import ( softcap_inplace_logits as fused_softcap, ) from sglang.srt.distributed.device_communicators import triton_symm_mem_ag -from sglang.srt.environ import envs from sglang.srt.layers.dp_attention import ( DpPaddingMode, attn_tp_all_gather, @@ -36,11 +35,9 @@ from sglang.srt.layers.dp_attention import ( get_dp_dtype, get_dp_hidden_size, ) -from sglang.srt.layers.utils.logprob import ( - InputLogprobsResult, - get_token_ids_logprobs_chunk, +from sglang.srt.layers.logprob_processor import ( + InputLogprobProcessor, get_token_ids_logprobs_prefill, - get_top_logprobs_chunk, get_top_logprobs_prefill, ) from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding @@ -373,10 +370,7 @@ class LogitsProcessor(nn.Module): skip_entry_sync=True, ) - # enable chunked logprobs processing - self.enable_logprobs_chunk = envs.SGLANG_ENABLE_LOGITS_PROCESSER_CHUNK.get() - # chunk size for logprobs processing - self.logprobs_chunk_size = envs.SGLANG_LOGITS_PROCESSER_CHUNK_SIZE.get() + self.input_logprob_processor = InputLogprobProcessor() def forward( self, @@ -457,38 +451,17 @@ class LogitsProcessor(nn.Module): mm_input_embeds=logits_metadata.mm_input_embeds, ) - # Start to process input logprobs - # Determine whether to use chunked or non-chunked logits processing. - # Skip chunking if: - # 1. Chunking is disabled - # 2. Total count is below chunk size threshold - # 3. DP attention all-gather is enabled (can use "enable_dp_lm_head" to enable chunking) - should_skip_chunking = ( - not self.enable_logprobs_chunk - or pruned_states.shape[0] <= self.logprobs_chunk_size - or self.do_tensor_parallel_all_gather_dp_attn + logprobs_result, sampled_logits = self.input_logprob_processor.forward( + pruned_states=pruned_states, + sample_indices=sample_indices, + input_logprob_indices=input_logprob_indices, + token_to_seq_idx=token_to_seq_idx, + lm_head=lm_head, + get_logits_fn=self._get_logits, + logits_metadata=logits_metadata, + skip_chunking_for_dp_attn=self.do_tensor_parallel_all_gather_dp_attn, ) - if should_skip_chunking: - # Compute logits for both input and sampled tokens. - logits = self._get_logits(pruned_states, lm_head, logits_metadata) - sampled_logits = ( - logits[sample_indices] if sample_indices is not None else logits - ) - input_logits = logits[input_logprob_indices] - del logits - - logprobs_result = self.process_input_logprobs(input_logits, logits_metadata) - else: - logprobs_result, sampled_logits = self.process_input_logprobs_by_chunk( - pruned_states, - sample_indices, - input_logprob_indices, - token_to_seq_idx, - lm_head, - logits_metadata, - ) - return LogitsProcessorOutput( next_token_logits=sampled_logits, hidden_states=hidden_states_to_store, @@ -696,204 +669,6 @@ class LogitsProcessor(nn.Module): return hidden_states_to_store - def process_input_logprobs(self, input_logits, logits_metadata: LogitsMetadata): - input_logprobs = torch.nn.functional.log_softmax(input_logits, dim=-1) - - # Get the logprob of top-k tokens - if logits_metadata.extend_return_top_logprob: - ( - input_top_logprobs_val, - input_top_logprobs_idx, - ) = get_top_logprobs_prefill(input_logprobs, logits_metadata) - else: - input_top_logprobs_val = input_top_logprobs_idx = None - - # Get the logprob of given token id - if logits_metadata.extend_token_ids_logprob: - ( - input_token_ids_logprobs_val, - input_token_ids_logprobs_idx, - ) = get_token_ids_logprobs_prefill(input_logprobs, logits_metadata) - else: - input_token_ids_logprobs_val = input_token_ids_logprobs_idx = None - - input_token_logprobs = input_logprobs[ - torch.arange(input_logprobs.shape[0], device=input_logprobs.device), - logits_metadata.extend_input_logprob_token_ids_gpu, - ] - - return InputLogprobsResult( - input_token_logprobs=input_token_logprobs, - input_top_logprobs_val=input_top_logprobs_val, - input_top_logprobs_idx=input_top_logprobs_idx, - input_token_ids_logprobs_val=input_token_ids_logprobs_val, - input_token_ids_logprobs_idx=input_token_ids_logprobs_idx, - ) - - def process_input_logprobs_by_chunk( - self, - pruned_states: torch.Tensor, - sample_indices: torch.Tensor, - input_logprob_indices: torch.Tensor, - token_to_seq_idx: list[int], - lm_head: VocabParallelEmbedding, - logits_metadata: LogitsMetadata, - ) -> Tuple[InputLogprobsResult, torch.Tensor]: - """ - compute logprobs for the output token from the hidden states. - To avoid using too much memory, we split pruned_states into chunks of - rows to compute input_logprobs separately, then concatenate the results. - - Returns: - InputLogprobsResult: logprobs result - torch.Tensor: sampled logits - """ - - # The peak memory usage is proportional to the chunk size. - chunk_size = self.logprobs_chunk_size - total_size = pruned_states.shape[0] - num_chunks = (total_size + chunk_size - 1) // chunk_size - - input_token_logprobs = [] - if logits_metadata.extend_return_top_logprob: - input_top_logprobs_val = [] - input_top_logprobs_idx = [] - else: - input_top_logprobs_val = None - input_top_logprobs_idx = None - if logits_metadata.extend_token_ids_logprob: - input_token_ids_logprobs_val = [] - input_token_ids_logprobs_idx = [] - else: - input_token_ids_logprobs_val = None - input_token_ids_logprobs_idx = None - - # If a single sequence is split into multiple chunks, we need to keep track - # of the pruned length of the sequences in the previous chunks. - split_len_topk = 0 - split_len_token_ids = 0 - - for i in range(num_chunks): - start_idx = i * chunk_size - end_idx = min((i + 1) * chunk_size, total_size) - - # Notify lm_head LoRA about the current chunk so it can swap - # to the precomputed per-chunk batch_info. This is a no-op - # for non-LoRA lm_head modules. - if hasattr(lm_head, "set_lm_head_pass"): - lm_head.set_lm_head_pass(i) - - # Get indices for this chunk - chunk_mask = (input_logprob_indices >= start_idx) & ( - input_logprob_indices < end_idx - ) - global_indices = input_logprob_indices[chunk_mask] - chunk_indices = global_indices - start_idx - # Get the positions in the original array where chunk_mask is True - # This is needed to correctly index into extend_input_logprob_token_ids_gpu - mask_indices = torch.nonzero(chunk_mask, as_tuple=True)[0] - - # Get the logits for this chunk. Each chunk must own its output: - # writing through the shared graph logits buffer would alias - # chunks whose shape happens to match the buffer. - chunk_states = pruned_states[start_idx:end_idx] - chunk_logits = self._get_logits( - chunk_states, lm_head, logits_metadata, use_logits_buffer=False - ) - - # Initialize sampled_logits on first chunk - if i == 0: - sampled_logits = torch.empty( - (sample_indices.shape[0], chunk_logits.shape[1]), - dtype=chunk_logits.dtype, - device=chunk_logits.device, - ) - - # Handle sampled logits for the chunk if needed - # This must be done before the continue statement to ensure all sampled_logits are filled - chunk_sample_mask = (sample_indices >= start_idx) & ( - sample_indices < end_idx - ) - if chunk_sample_mask.any(): - chunk_sample_indices = sample_indices[chunk_sample_mask] - start_idx - sampled_logits[chunk_sample_mask] = chunk_logits[chunk_sample_indices] - - # If there are no input logprobs in this chunk, skip the rest - if chunk_indices.numel() == 0: - continue - - # Compute the logprobs of the chunk - chunk_input_logprobs = chunk_logits[chunk_indices] - chunk_input_logprobs = torch.nn.functional.log_softmax( - chunk_input_logprobs, dim=-1 - ) - - # For each chunk, we need to get the slice of the token_to_seq_idx - chunk_slice = slice( - token_to_seq_idx[start_idx], token_to_seq_idx[end_idx] + 1 - ) - - # Get the logprob of top-k tokens - if logits_metadata.extend_return_top_logprob: - top_k_nums = logits_metadata.top_logprobs_nums[chunk_slice] - pruned_lens = logits_metadata.extend_logprob_pruned_lens_cpu[ - chunk_slice - ] - split_len_topk = get_top_logprobs_chunk( - chunk_input_logprobs, - logits_metadata, - top_k_nums, - pruned_lens, - input_top_logprobs_val, - input_top_logprobs_idx, - split_len_topk, - ) - - # Get the logprob of given token id - if logits_metadata.extend_token_ids_logprob: - token_ids_logprobs = logits_metadata.token_ids_logprobs[chunk_slice] - pruned_lens = logits_metadata.extend_logprob_pruned_lens_cpu[ - chunk_slice - ] - split_len_token_ids = get_token_ids_logprobs_chunk( - chunk_input_logprobs, - token_ids_logprobs, - pruned_lens, - input_token_ids_logprobs_val, - input_token_ids_logprobs_idx, - split_len_token_ids, - ) - - # Get the logprob of the requested token ids - chunk_input_token_logprobs = chunk_input_logprobs[ - torch.arange( - chunk_input_logprobs.shape[0], device=chunk_input_logprobs.device - ), - logits_metadata.extend_input_logprob_token_ids_gpu[mask_indices], - ] - input_token_logprobs.append(chunk_input_token_logprobs) - - # Restore the full-pruned lm_head batch_info after chunk iteration. - if hasattr(lm_head, "reset_lm_head_pass"): - assert hasattr( - lm_head, "set_lm_head_pass" - ), "lm_head must have set_lm_head_pass method and reset_lm_head_pass method at the same time" - lm_head.reset_lm_head_pass() - - # Concatenate the results - input_token_logprobs = torch.cat(input_token_logprobs, dim=0) - - return ( - InputLogprobsResult( - input_token_logprobs=input_token_logprobs, - input_top_logprobs_val=input_top_logprobs_val, - input_top_logprobs_idx=input_top_logprobs_idx, - input_token_ids_logprobs_val=input_token_ids_logprobs_val, - input_token_ids_logprobs_idx=input_token_ids_logprobs_idx, - ), - sampled_logits, - ) - def _get_logits( self, hidden_states: torch.Tensor, diff --git a/python/sglang/srt/layers/utils/logprob.py b/python/sglang/srt/layers/logprob_processor.py similarity index 51% rename from python/sglang/srt/layers/utils/logprob.py rename to python/sglang/srt/layers/logprob_processor.py index fde20e211..0b7833974 100644 --- a/python/sglang/srt/layers/utils/logprob.py +++ b/python/sglang/srt/layers/logprob_processor.py @@ -2,7 +2,7 @@ from __future__ import annotations import dataclasses from enum import Enum, auto -from typing import TYPE_CHECKING, List, Optional +from typing import TYPE_CHECKING, Callable, List, Optional, Tuple import torch @@ -10,6 +10,7 @@ from sglang.srt.environ import envs if TYPE_CHECKING: from sglang.srt.layers.logits_processor import LogitsMetadata + from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding class LogprobStage(Enum): @@ -355,3 +356,263 @@ def compute_spec_v2_logprobs( ) = get_token_ids_logprobs( gathered_logprobs, token_ids_logprobs_expanded, no_copy_to_cpu=True ) + + +class InputLogprobProcessor: + """Input (prefill) logprob processing: single-pass or chunked. + + Logits are computed through the injected ``get_logits_fn(hidden_states, + lm_head, logits_metadata)`` callable, so this class stays decoupled from + the lm_head / TP-gather machinery in LogitsProcessor. + """ + + def __init__(self): + # enable chunked logprobs processing + self.enable_logprobs_chunk = envs.SGLANG_ENABLE_LOGITS_PROCESSER_CHUNK.get() + # chunk size for logprobs processing + self.logprobs_chunk_size = envs.SGLANG_LOGITS_PROCESSER_CHUNK_SIZE.get() + + def forward( + self, + pruned_states: torch.Tensor, + sample_indices: Optional[torch.Tensor], + input_logprob_indices: torch.Tensor, + token_to_seq_idx: list[int], + lm_head: VocabParallelEmbedding, + get_logits_fn: Callable, + logits_metadata: LogitsMetadata, + skip_chunking_for_dp_attn: bool = False, + ) -> Tuple[InputLogprobsResult, torch.Tensor]: + # Start to process input logprobs + # Determine whether to use chunked or non-chunked logits processing. + # Skip chunking if: + # 1. Chunking is disabled + # 2. Total count is below chunk size threshold + # 3. DP attention all-gather is enabled (can use "enable_dp_lm_head" to enable chunking) + should_skip_chunking = ( + not self.enable_logprobs_chunk + or pruned_states.shape[0] <= self.logprobs_chunk_size + or skip_chunking_for_dp_attn + ) + + if should_skip_chunking: + # Compute logits for both input and sampled tokens. + logits = get_logits_fn(pruned_states, lm_head, logits_metadata) + sampled_logits = ( + logits[sample_indices] if sample_indices is not None else logits + ) + input_logits = logits[input_logprob_indices] + del logits + + logprobs_result = self.process_input_logprobs(input_logits, logits_metadata) + else: + logprobs_result, sampled_logits = self.process_input_logprobs_by_chunk( + pruned_states, + sample_indices, + input_logprob_indices, + token_to_seq_idx, + lm_head, + get_logits_fn, + logits_metadata, + ) + + return logprobs_result, sampled_logits + + def process_input_logprobs(self, input_logits, logits_metadata: LogitsMetadata): + input_logprobs = torch.nn.functional.log_softmax(input_logits, dim=-1) + + # Get the logprob of top-k tokens + if logits_metadata.extend_return_top_logprob: + ( + input_top_logprobs_val, + input_top_logprobs_idx, + ) = get_top_logprobs_prefill(input_logprobs, logits_metadata) + else: + input_top_logprobs_val = input_top_logprobs_idx = None + + # Get the logprob of given token id + if logits_metadata.extend_token_ids_logprob: + ( + input_token_ids_logprobs_val, + input_token_ids_logprobs_idx, + ) = get_token_ids_logprobs_prefill(input_logprobs, logits_metadata) + else: + input_token_ids_logprobs_val = input_token_ids_logprobs_idx = None + + input_token_logprobs = input_logprobs[ + torch.arange(input_logprobs.shape[0], device=input_logprobs.device), + logits_metadata.extend_input_logprob_token_ids_gpu, + ] + + return InputLogprobsResult( + input_token_logprobs=input_token_logprobs, + input_top_logprobs_val=input_top_logprobs_val, + input_top_logprobs_idx=input_top_logprobs_idx, + input_token_ids_logprobs_val=input_token_ids_logprobs_val, + input_token_ids_logprobs_idx=input_token_ids_logprobs_idx, + ) + + def process_input_logprobs_by_chunk( + self, + pruned_states: torch.Tensor, + sample_indices: torch.Tensor, + input_logprob_indices: torch.Tensor, + token_to_seq_idx: list[int], + lm_head: VocabParallelEmbedding, + get_logits_fn: Callable, + logits_metadata: LogitsMetadata, + ) -> Tuple[InputLogprobsResult, torch.Tensor]: + """ + compute logprobs for the output token from the hidden states. + To avoid using too much memory, we split pruned_states into chunks of + rows to compute input_logprobs separately, then concatenate the results. + + Returns: + InputLogprobsResult: logprobs result + torch.Tensor: sampled logits + """ + + # The peak memory usage is proportional to the chunk size. + chunk_size = self.logprobs_chunk_size + total_size = pruned_states.shape[0] + num_chunks = (total_size + chunk_size - 1) // chunk_size + + input_token_logprobs = [] + if logits_metadata.extend_return_top_logprob: + input_top_logprobs_val = [] + input_top_logprobs_idx = [] + else: + input_top_logprobs_val = None + input_top_logprobs_idx = None + if logits_metadata.extend_token_ids_logprob: + input_token_ids_logprobs_val = [] + input_token_ids_logprobs_idx = [] + else: + input_token_ids_logprobs_val = None + input_token_ids_logprobs_idx = None + + # If a single sequence is split into multiple chunks, we need to keep track + # of the pruned length of the sequences in the previous chunks. + split_len_topk = 0 + split_len_token_ids = 0 + + for i in range(num_chunks): + start_idx = i * chunk_size + end_idx = min((i + 1) * chunk_size, total_size) + + # Notify lm_head LoRA about the current chunk so it can swap + # to the precomputed per-chunk batch_info. This is a no-op + # for non-LoRA lm_head modules. + if hasattr(lm_head, "set_lm_head_pass"): + lm_head.set_lm_head_pass(i) + + # Get indices for this chunk + chunk_mask = (input_logprob_indices >= start_idx) & ( + input_logprob_indices < end_idx + ) + global_indices = input_logprob_indices[chunk_mask] + chunk_indices = global_indices - start_idx + # Get the positions in the original array where chunk_mask is True + # This is needed to correctly index into extend_input_logprob_token_ids_gpu + mask_indices = torch.nonzero(chunk_mask, as_tuple=True)[0] + + # Get the logits for this chunk. Each chunk must own its output: + # writing through the shared graph logits buffer would alias + # chunks whose shape happens to match the buffer. + chunk_states = pruned_states[start_idx:end_idx] + chunk_logits = get_logits_fn( + chunk_states, lm_head, logits_metadata, use_logits_buffer=False + ) + + # Initialize sampled_logits on first chunk + if i == 0: + sampled_logits = torch.empty( + (sample_indices.shape[0], chunk_logits.shape[1]), + dtype=chunk_logits.dtype, + device=chunk_logits.device, + ) + + # Handle sampled logits for the chunk if needed + # This must be done before the continue statement to ensure all sampled_logits are filled + chunk_sample_mask = (sample_indices >= start_idx) & ( + sample_indices < end_idx + ) + if chunk_sample_mask.any(): + chunk_sample_indices = sample_indices[chunk_sample_mask] - start_idx + sampled_logits[chunk_sample_mask] = chunk_logits[chunk_sample_indices] + + # If there are no input logprobs in this chunk, skip the rest + if chunk_indices.numel() == 0: + continue + + # Compute the logprobs of the chunk + chunk_input_logprobs = chunk_logits[chunk_indices] + chunk_input_logprobs = torch.nn.functional.log_softmax( + chunk_input_logprobs, dim=-1 + ) + + # For each chunk, we need to get the slice of the token_to_seq_idx + chunk_slice = slice( + token_to_seq_idx[start_idx], token_to_seq_idx[end_idx] + 1 + ) + + # Get the logprob of top-k tokens + if logits_metadata.extend_return_top_logprob: + top_k_nums = logits_metadata.top_logprobs_nums[chunk_slice] + pruned_lens = logits_metadata.extend_logprob_pruned_lens_cpu[ + chunk_slice + ] + split_len_topk = get_top_logprobs_chunk( + chunk_input_logprobs, + logits_metadata, + top_k_nums, + pruned_lens, + input_top_logprobs_val, + input_top_logprobs_idx, + split_len_topk, + ) + + # Get the logprob of given token id + if logits_metadata.extend_token_ids_logprob: + token_ids_logprobs = logits_metadata.token_ids_logprobs[chunk_slice] + pruned_lens = logits_metadata.extend_logprob_pruned_lens_cpu[ + chunk_slice + ] + split_len_token_ids = get_token_ids_logprobs_chunk( + chunk_input_logprobs, + token_ids_logprobs, + pruned_lens, + input_token_ids_logprobs_val, + input_token_ids_logprobs_idx, + split_len_token_ids, + ) + + # Get the logprob of the requested token ids + chunk_input_token_logprobs = chunk_input_logprobs[ + torch.arange( + chunk_input_logprobs.shape[0], device=chunk_input_logprobs.device + ), + logits_metadata.extend_input_logprob_token_ids_gpu[mask_indices], + ] + input_token_logprobs.append(chunk_input_token_logprobs) + + # Restore the full-pruned lm_head batch_info after chunk iteration. + if hasattr(lm_head, "reset_lm_head_pass"): + assert hasattr( + lm_head, "set_lm_head_pass" + ), "lm_head must have set_lm_head_pass method and reset_lm_head_pass method at the same time" + lm_head.reset_lm_head_pass() + + # Concatenate the results + input_token_logprobs = torch.cat(input_token_logprobs, dim=0) + + return ( + InputLogprobsResult( + input_token_logprobs=input_token_logprobs, + input_top_logprobs_val=input_top_logprobs_val, + input_top_logprobs_idx=input_top_logprobs_idx, + input_token_ids_logprobs_val=input_token_ids_logprobs_val, + input_token_ids_logprobs_idx=input_token_ids_logprobs_idx, + ), + sampled_logits, + ) diff --git a/python/sglang/srt/layers/sampler.py b/python/sglang/srt/layers/sampler.py index dcbdc27ed..68280bfd3 100644 --- a/python/sglang/srt/layers/sampler.py +++ b/python/sglang/srt/layers/sampler.py @@ -11,7 +11,7 @@ from sglang.srt.layers.dp_attention import ( is_dp_attention_enabled, ) from sglang.srt.layers.logits_processor import LogitsProcessorOutput -from sglang.srt.layers.utils.logprob import get_token_ids_logprobs, get_top_logprobs +from sglang.srt.layers.logprob_processor import get_token_ids_logprobs, get_top_logprobs from sglang.srt.runtime_context import get_parallel, get_server_args from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo from sglang.srt.sampling.sampling_params import TOP_K_ALL diff --git a/python/sglang/srt/lora/layers.py b/python/sglang/srt/lora/layers.py index f0032b80b..613342a47 100644 --- a/python/sglang/srt/lora/layers.py +++ b/python/sglang/srt/lora/layers.py @@ -390,7 +390,7 @@ class ParallelLMHeadWithLoRA(BaseLayerWithLoRA): def set_lm_head_pass(self, pass_idx: int): """Set the active lm_head pass index before a logprobs chunk. - Called by LogitsProcessor.process_input_logprobs_by_chunk() before + Called by InputLogprobProcessor.process_input_logprobs_by_chunk() before each chunk's _get_logits call. _get_lm_head_batch_info() will resolve to lm_head_pass_batch_infos[pass_idx]. """ diff --git a/python/sglang/srt/lora/utils.py b/python/sglang/srt/lora/utils.py index e8301d5b4..619c65df2 100644 --- a/python/sglang/srt/lora/utils.py +++ b/python/sglang/srt/lora/utils.py @@ -544,7 +544,7 @@ def build_lm_head_pass_segments( """ Precompute per-pass segment info for lm_head LoRA logprobs processing. - When LogitsProcessor uses chunked logprobs processing + When InputLogprobProcessor uses chunked logprobs processing (process_input_logprobs_by_chunk), pruned hidden states are split into fixed-size passes. Each pass needs its own segmentation (weight_indices, seg_lens) so that lm_head LoRA operates on the diff --git a/python/sglang/srt/speculative/eagle_worker_common.py b/python/sglang/srt/speculative/eagle_worker_common.py index 1b37b311a..146ad94ed 100644 --- a/python/sglang/srt/speculative/eagle_worker_common.py +++ b/python/sglang/srt/speculative/eagle_worker_common.py @@ -8,7 +8,7 @@ from sglang.kernels.ops.speculative.cache_locs import ( assign_draft_cache_locs_contiguous, ) from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func -from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs +from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs from sglang.srt.managers.utils import GenerationBatchResult from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index 83dc3d535..01ca29f55 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -9,7 +9,7 @@ from sglang.kernels.ops.speculative.cache_locs import ( assign_extend_cache_locs_func as assign_extend_cache_locs_func, ) from sglang.srt.distributed.parallel_state_wrapper import ParallelState -from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs +from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.tp_worker import TpModelWorker