refactor logprob processor layer (#20071)
This commit is contained in:
@@ -25,7 +25,6 @@ from sglang.kernels.ops.activation.softcap import (
|
|||||||
softcap_inplace_logits as fused_softcap,
|
softcap_inplace_logits as fused_softcap,
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed.device_communicators import triton_symm_mem_ag
|
from sglang.srt.distributed.device_communicators import triton_symm_mem_ag
|
||||||
from sglang.srt.environ import envs
|
|
||||||
from sglang.srt.layers.dp_attention import (
|
from sglang.srt.layers.dp_attention import (
|
||||||
DpPaddingMode,
|
DpPaddingMode,
|
||||||
attn_tp_all_gather,
|
attn_tp_all_gather,
|
||||||
@@ -36,11 +35,9 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
get_dp_dtype,
|
get_dp_dtype,
|
||||||
get_dp_hidden_size,
|
get_dp_hidden_size,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.utils.logprob import (
|
from sglang.srt.layers.logprob_processor import (
|
||||||
InputLogprobsResult,
|
InputLogprobProcessor,
|
||||||
get_token_ids_logprobs_chunk,
|
|
||||||
get_token_ids_logprobs_prefill,
|
get_token_ids_logprobs_prefill,
|
||||||
get_top_logprobs_chunk,
|
|
||||||
get_top_logprobs_prefill,
|
get_top_logprobs_prefill,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||||
@@ -373,10 +370,7 @@ class LogitsProcessor(nn.Module):
|
|||||||
skip_entry_sync=True,
|
skip_entry_sync=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
# enable chunked logprobs processing
|
self.input_logprob_processor = InputLogprobProcessor()
|
||||||
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(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -457,36 +451,15 @@ class LogitsProcessor(nn.Module):
|
|||||||
mm_input_embeds=logits_metadata.mm_input_embeds,
|
mm_input_embeds=logits_metadata.mm_input_embeds,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Start to process input logprobs
|
logprobs_result, sampled_logits = self.input_logprob_processor.forward(
|
||||||
# Determine whether to use chunked or non-chunked logits processing.
|
pruned_states=pruned_states,
|
||||||
# Skip chunking if:
|
sample_indices=sample_indices,
|
||||||
# 1. Chunking is disabled
|
input_logprob_indices=input_logprob_indices,
|
||||||
# 2. Total count is below chunk size threshold
|
token_to_seq_idx=token_to_seq_idx,
|
||||||
# 3. DP attention all-gather is enabled (can use "enable_dp_lm_head" to enable chunking)
|
lm_head=lm_head,
|
||||||
should_skip_chunking = (
|
get_logits_fn=self._get_logits,
|
||||||
not self.enable_logprobs_chunk
|
logits_metadata=logits_metadata,
|
||||||
or pruned_states.shape[0] <= self.logprobs_chunk_size
|
skip_chunking_for_dp_attn=self.do_tensor_parallel_all_gather_dp_attn,
|
||||||
or 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(
|
return LogitsProcessorOutput(
|
||||||
@@ -696,204 +669,6 @@ class LogitsProcessor(nn.Module):
|
|||||||
|
|
||||||
return hidden_states_to_store
|
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(
|
def _get_logits(
|
||||||
self,
|
self,
|
||||||
hidden_states: torch.Tensor,
|
hidden_states: torch.Tensor,
|
||||||
|
|||||||
+262
-1
@@ -2,7 +2,7 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import dataclasses
|
import dataclasses
|
||||||
from enum import Enum, auto
|
from enum import Enum, auto
|
||||||
from typing import TYPE_CHECKING, List, Optional
|
from typing import TYPE_CHECKING, Callable, List, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -10,6 +10,7 @@ from sglang.srt.environ import envs
|
|||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
from sglang.srt.layers.logits_processor import LogitsMetadata
|
from sglang.srt.layers.logits_processor import LogitsMetadata
|
||||||
|
from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||||
|
|
||||||
|
|
||||||
class LogprobStage(Enum):
|
class LogprobStage(Enum):
|
||||||
@@ -355,3 +356,263 @@ def compute_spec_v2_logprobs(
|
|||||||
) = get_token_ids_logprobs(
|
) = get_token_ids_logprobs(
|
||||||
gathered_logprobs, token_ids_logprobs_expanded, no_copy_to_cpu=True
|
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,
|
||||||
|
)
|
||||||
@@ -11,7 +11,7 @@ from sglang.srt.layers.dp_attention import (
|
|||||||
is_dp_attention_enabled,
|
is_dp_attention_enabled,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
|
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.runtime_context import get_parallel, get_server_args
|
||||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||||
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
from sglang.srt.sampling.sampling_params import TOP_K_ALL
|
||||||
|
|||||||
@@ -390,7 +390,7 @@ class ParallelLMHeadWithLoRA(BaseLayerWithLoRA):
|
|||||||
def set_lm_head_pass(self, pass_idx: int):
|
def set_lm_head_pass(self, pass_idx: int):
|
||||||
"""Set the active lm_head pass index before a logprobs chunk.
|
"""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
|
each chunk's _get_logits call. _get_lm_head_batch_info() will
|
||||||
resolve to lm_head_pass_batch_infos[pass_idx].
|
resolve to lm_head_pass_batch_infos[pass_idx].
|
||||||
"""
|
"""
|
||||||
|
|||||||
@@ -544,7 +544,7 @@ def build_lm_head_pass_segments(
|
|||||||
"""
|
"""
|
||||||
Precompute per-pass segment info for lm_head LoRA logprobs processing.
|
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
|
(process_input_logprobs_by_chunk), pruned hidden states are split into
|
||||||
fixed-size passes. Each pass needs its own segmentation
|
fixed-size passes. Each pass needs its own segmentation
|
||||||
(weight_indices, seg_lens) so that lm_head LoRA operates on the
|
(weight_indices, seg_lens) so that lm_head LoRA operates on the
|
||||||
|
|||||||
@@ -8,7 +8,7 @@ from sglang.kernels.ops.speculative.cache_locs import (
|
|||||||
assign_draft_cache_locs_contiguous,
|
assign_draft_cache_locs_contiguous,
|
||||||
)
|
)
|
||||||
from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func
|
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.managers.utils import GenerationBatchResult
|
||||||
from sglang.srt.model_executor.forward_batch_info import (
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
CaptureHiddenMode,
|
CaptureHiddenMode,
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ from sglang.kernels.ops.speculative.cache_locs import (
|
|||||||
assign_extend_cache_locs_func as assign_extend_cache_locs_func,
|
assign_extend_cache_locs_func as assign_extend_cache_locs_func,
|
||||||
)
|
)
|
||||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
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.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
|
|||||||
Reference in New Issue
Block a user