Remove dead fields and always-False plumbing across SB / FB / LogitsMetadata (#26551)
This commit is contained in:
@@ -754,9 +754,7 @@ class TboForwardBatchPreparer:
|
|||||||
global_num_tokens_for_logprob_cpu=None,
|
global_num_tokens_for_logprob_cpu=None,
|
||||||
sampling_info=None,
|
sampling_info=None,
|
||||||
# For logits and logprobs post processing, thus we do not care
|
# For logits and logprobs post processing, thus we do not care
|
||||||
temp_scaled_logprobs=False,
|
|
||||||
temperature=None,
|
temperature=None,
|
||||||
top_p_normalized_logprobs=False,
|
|
||||||
top_p=None,
|
top_p=None,
|
||||||
mm_inputs=None,
|
mm_inputs=None,
|
||||||
top_logprobs_nums=None,
|
top_logprobs_nums=None,
|
||||||
|
|||||||
@@ -43,7 +43,6 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.layers.utils.logprob import (
|
from sglang.srt.layers.utils.logprob import (
|
||||||
InputLogprobsResult,
|
InputLogprobsResult,
|
||||||
compute_temp_top_p_normalized_logprobs,
|
|
||||||
get_token_ids_logprobs_chunk,
|
get_token_ids_logprobs_chunk,
|
||||||
get_token_ids_logprobs_prefill,
|
get_token_ids_logprobs_prefill,
|
||||||
get_top_logprobs_chunk,
|
get_top_logprobs_chunk,
|
||||||
@@ -132,9 +131,7 @@ class LogitsMetadata:
|
|||||||
token_ids_logprobs: Optional[List[List[int]]] = None
|
token_ids_logprobs: Optional[List[List[int]]] = None
|
||||||
|
|
||||||
# logits and logprobs post processing
|
# logits and logprobs post processing
|
||||||
temp_scaled_logprobs: bool = False
|
|
||||||
temperature: torch.Tensor = None
|
temperature: torch.Tensor = None
|
||||||
top_p_normalized_logprobs: bool = False
|
|
||||||
top_p: torch.Tensor = None
|
top_p: torch.Tensor = None
|
||||||
|
|
||||||
# DP attention metadata. Not needed when DP attention is not used.
|
# DP attention metadata. Not needed when DP attention is not used.
|
||||||
@@ -362,9 +359,6 @@ class LogitsProcessor(nn.Module):
|
|||||||
)
|
)
|
||||||
|
|
||||||
# Start to process input logprobs
|
# Start to process input logprobs
|
||||||
# Normalize the logprob w/o temperature, top-p
|
|
||||||
self._expand_metadata_for_logprobs(logits_metadata, pruned_states.device)
|
|
||||||
|
|
||||||
# Determine whether to use chunked or non-chunked logits processing.
|
# Determine whether to use chunked or non-chunked logits processing.
|
||||||
# Skip chunking if:
|
# Skip chunking if:
|
||||||
# 1. Chunking is disabled
|
# 1. Chunking is disabled
|
||||||
@@ -617,29 +611,8 @@ class LogitsProcessor(nn.Module):
|
|||||||
|
|
||||||
return hidden_states_to_store
|
return hidden_states_to_store
|
||||||
|
|
||||||
def _expand_metadata_for_logprobs(
|
|
||||||
self, logits_metadata: LogitsMetadata, device: torch.device
|
|
||||||
):
|
|
||||||
pruned_lens = torch.tensor(
|
|
||||||
logits_metadata.extend_logprob_pruned_lens_cpu,
|
|
||||||
dtype=torch.int64,
|
|
||||||
pin_memory=is_pin_memory_available(),
|
|
||||||
).to(device, non_blocking=True)
|
|
||||||
if logits_metadata.temp_scaled_logprobs:
|
|
||||||
logits_metadata.temperature = torch.repeat_interleave(
|
|
||||||
logits_metadata.temperature.view(-1),
|
|
||||||
pruned_lens,
|
|
||||||
).view(-1, 1)
|
|
||||||
if logits_metadata.top_p_normalized_logprobs:
|
|
||||||
logits_metadata.top_p = torch.repeat_interleave(
|
|
||||||
logits_metadata.top_p,
|
|
||||||
pruned_lens,
|
|
||||||
)
|
|
||||||
|
|
||||||
def process_input_logprobs(self, input_logits, logits_metadata: LogitsMetadata):
|
def process_input_logprobs(self, input_logits, logits_metadata: LogitsMetadata):
|
||||||
input_logprobs = compute_temp_top_p_normalized_logprobs(
|
input_logprobs = torch.nn.functional.log_softmax(input_logits, dim=-1)
|
||||||
input_logits, logits_metadata
|
|
||||||
)
|
|
||||||
|
|
||||||
# Get the logprob of top-k tokens
|
# Get the logprob of top-k tokens
|
||||||
if logits_metadata.extend_return_top_logprob:
|
if logits_metadata.extend_return_top_logprob:
|
||||||
@@ -762,26 +735,8 @@ class LogitsProcessor(nn.Module):
|
|||||||
|
|
||||||
# Compute the logprobs of the chunk
|
# Compute the logprobs of the chunk
|
||||||
chunk_input_logprobs = chunk_logits[chunk_indices]
|
chunk_input_logprobs = chunk_logits[chunk_indices]
|
||||||
# Only index per-token arrays when the corresponding feature is active.
|
chunk_input_logprobs = torch.nn.functional.log_softmax(
|
||||||
# Otherwise these tensors can be per-sequence (or scalars), which can
|
chunk_input_logprobs, dim=-1
|
||||||
# cause out-of-bounds indexing on GPU.
|
|
||||||
chunk_temperature = (
|
|
||||||
logits_metadata.temperature[global_indices]
|
|
||||||
if logits_metadata.temp_scaled_logprobs
|
|
||||||
and logits_metadata.temperature is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
chunk_top_p = (
|
|
||||||
logits_metadata.top_p[global_indices]
|
|
||||||
if logits_metadata.top_p_normalized_logprobs
|
|
||||||
and logits_metadata.top_p is not None
|
|
||||||
else None
|
|
||||||
)
|
|
||||||
chunk_input_logprobs = compute_temp_top_p_normalized_logprobs(
|
|
||||||
chunk_input_logprobs,
|
|
||||||
logits_metadata,
|
|
||||||
chunk_top_p,
|
|
||||||
chunk_temperature,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
# For each chunk, we need to get the slice of the token_to_seq_idx
|
# For each chunk, we need to get the slice of the token_to_seq_idx
|
||||||
|
|||||||
@@ -29,40 +29,6 @@ class InputLogprobsResult:
|
|||||||
input_token_ids_logprobs_idx: Optional[List] = None
|
input_token_ids_logprobs_idx: Optional[List] = None
|
||||||
|
|
||||||
|
|
||||||
def compute_temp_top_p_normalized_logprobs(
|
|
||||||
last_logits: torch.Tensor,
|
|
||||||
logits_metadata: LogitsMetadata,
|
|
||||||
top_p: Optional[torch.Tensor] = None,
|
|
||||||
temperature: Optional[torch.Tensor] = None,
|
|
||||||
) -> torch.Tensor:
|
|
||||||
"""
|
|
||||||
compute logprobs for the output token from the given logits.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
torch.Tensor: logprobs from logits
|
|
||||||
"""
|
|
||||||
if top_p is None:
|
|
||||||
top_p = logits_metadata.top_p
|
|
||||||
if temperature is None:
|
|
||||||
temperature = logits_metadata.temperature
|
|
||||||
|
|
||||||
# Scale logits if temperature scaling is enabled
|
|
||||||
if logits_metadata.temp_scaled_logprobs:
|
|
||||||
last_logits = last_logits / temperature
|
|
||||||
|
|
||||||
# Normalize logprobs if top_p normalization is enabled
|
|
||||||
# NOTE: only normalize logprobs when top_p is set and not equal to 1.0
|
|
||||||
if logits_metadata.top_p_normalized_logprobs and (top_p != 1.0).any():
|
|
||||||
from sglang.srt.layers.sampler import top_p_normalize_probs_torch
|
|
||||||
|
|
||||||
probs = torch.softmax(last_logits, dim=-1)
|
|
||||||
del last_logits
|
|
||||||
probs = top_p_normalize_probs_torch(probs, top_p)
|
|
||||||
return torch.log(probs)
|
|
||||||
else:
|
|
||||||
return torch.nn.functional.log_softmax(last_logits, dim=-1)
|
|
||||||
|
|
||||||
|
|
||||||
def get_top_logprobs_raw(
|
def get_top_logprobs_raw(
|
||||||
logprobs: torch.Tensor,
|
logprobs: torch.Tensor,
|
||||||
top_logprobs_nums: List[int],
|
top_logprobs_nums: List[int],
|
||||||
|
|||||||
@@ -1509,10 +1509,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
split_forward_count: int = 1
|
split_forward_count: int = 1
|
||||||
split_forward_batch: ForwardBatch = None
|
split_forward_batch: ForwardBatch = None
|
||||||
|
|
||||||
# For logits and logprob post processing (ForwardBatch keeps its own copies)
|
|
||||||
temp_scaled_logprobs: bool = False
|
|
||||||
top_p_normalized_logprobs: bool = False
|
|
||||||
|
|
||||||
# CPU mirror of req_pool_indices; schedule-path only (used in overlap_utils,
|
# CPU mirror of req_pool_indices; schedule-path only (used in overlap_utils,
|
||||||
# not read by ForwardBatch), stale in spec draft window
|
# not read by ForwardBatch), stale in spec draft window
|
||||||
req_pool_indices_cpu: torch.Tensor = None # shape: [b], int64
|
req_pool_indices_cpu: torch.Tensor = None # shape: [b], int64
|
||||||
@@ -1520,9 +1516,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# Forward-pass metrics
|
# Forward-pass metrics
|
||||||
fpm_start_time: float = 0.0
|
fpm_start_time: float = 0.0
|
||||||
|
|
||||||
# Stream
|
|
||||||
has_stream: bool = False
|
|
||||||
|
|
||||||
# Whether to return captured experts
|
# Whether to return captured experts
|
||||||
return_routed_experts: bool = False
|
return_routed_experts: bool = False
|
||||||
return_indexer_topk: bool = False
|
return_indexer_topk: bool = False
|
||||||
@@ -1673,7 +1666,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
model_config=model_config,
|
model_config=model_config,
|
||||||
enable_overlap=enable_overlap,
|
enable_overlap=enable_overlap,
|
||||||
return_logprob=return_logprob,
|
return_logprob=return_logprob,
|
||||||
has_stream=any(req.stream for req in reqs),
|
|
||||||
has_grammar=any(req.grammar for req in reqs),
|
has_grammar=any(req.grammar for req in reqs),
|
||||||
device=req_to_token_pool.device,
|
device=req_to_token_pool.device,
|
||||||
spec_algorithm=spec_algorithm,
|
spec_algorithm=spec_algorithm,
|
||||||
@@ -2585,7 +2577,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
self.top_logprobs_nums = None
|
self.top_logprobs_nums = None
|
||||||
self.token_ids_logprobs = None
|
self.token_ids_logprobs = None
|
||||||
|
|
||||||
self.has_stream = any(req.stream for req in self.reqs)
|
|
||||||
self.has_grammar = any(req.grammar for req in self.reqs)
|
self.has_grammar = any(req.grammar for req in self.reqs)
|
||||||
|
|
||||||
self.sampling_info.filter_batch(keep_indices, keep_indices_device)
|
self.sampling_info.filter_batch(keep_indices, keep_indices_device)
|
||||||
@@ -2641,7 +2632,6 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
self.multimodal_inputs.extend(other.multimodal_inputs)
|
self.multimodal_inputs.extend(other.multimodal_inputs)
|
||||||
|
|
||||||
self.return_logprob |= other.return_logprob
|
self.return_logprob |= other.return_logprob
|
||||||
self.has_stream |= other.has_stream
|
|
||||||
self.has_grammar |= other.has_grammar
|
self.has_grammar |= other.has_grammar
|
||||||
self.return_hidden_states |= other.return_hidden_states
|
self.return_hidden_states |= other.return_hidden_states
|
||||||
self.is_prefill_only = self.is_prefill_only and other.is_prefill_only
|
self.is_prefill_only = self.is_prefill_only and other.is_prefill_only
|
||||||
|
|||||||
@@ -406,9 +406,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
|
|||||||
# === Runtime-filled (set during the forward pass / cuda graph / managers; not at construction) ===
|
# === Runtime-filled (set during the forward pass / cuda graph / managers; not at construction) ===
|
||||||
# For logits and logprobs post processing
|
# For logits and logprobs post processing
|
||||||
next_token_logits_buffer: torch.Tensor = None
|
next_token_logits_buffer: torch.Tensor = None
|
||||||
temp_scaled_logprobs: bool = False
|
|
||||||
temperature: torch.Tensor = None
|
temperature: torch.Tensor = None
|
||||||
top_p_normalized_logprobs: bool = False
|
|
||||||
top_p: torch.Tensor = None
|
top_p: torch.Tensor = None
|
||||||
|
|
||||||
# For split prefill
|
# For split prefill
|
||||||
|
|||||||
@@ -771,9 +771,7 @@ class PiecewiseCudaGraphRunner:
|
|||||||
lora_ids=forward_batch.lora_ids,
|
lora_ids=forward_batch.lora_ids,
|
||||||
sampling_info=forward_batch.sampling_info,
|
sampling_info=forward_batch.sampling_info,
|
||||||
mm_inputs=forward_batch.mm_inputs,
|
mm_inputs=forward_batch.mm_inputs,
|
||||||
temp_scaled_logprobs=forward_batch.temp_scaled_logprobs,
|
|
||||||
temperature=forward_batch.temperature,
|
temperature=forward_batch.temperature,
|
||||||
top_p_normalized_logprobs=forward_batch.top_p_normalized_logprobs,
|
|
||||||
top_p=forward_batch.top_p,
|
top_p=forward_batch.top_p,
|
||||||
dimensions=forward_batch.dimensions,
|
dimensions=forward_batch.dimensions,
|
||||||
return_pooled_hidden_states=(
|
return_pooled_hidden_states=(
|
||||||
|
|||||||
@@ -181,7 +181,6 @@ class TestEAGLEServerAdditional(TestEAGLEServerBasic):
|
|||||||
"return_logprob": return_logprob,
|
"return_logprob": return_logprob,
|
||||||
"return_text_in_logprobs": True,
|
"return_text_in_logprobs": True,
|
||||||
"logprob_start_len": logprob_start_len,
|
"logprob_start_len": logprob_start_len,
|
||||||
"temp_scaled_logprobs": True,
|
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
return response.json()
|
return response.json()
|
||||||
|
|||||||
Reference in New Issue
Block a user