[Fix] Enable chunked input-logprob processing by default to cap peak memory (#31498)

This commit is contained in:
Liangsheng Yin
2026-07-17 13:19:03 -07:00
committed by GitHub
parent d389039337
commit 2c856abbe3
6 changed files with 101 additions and 12 deletions
+1 -1
View File
@@ -811,7 +811,7 @@ class Envs:
SGLANG_EMBEDDINGS_SPARSE_HEAD = EnvStr(None)
# Logits processor
SGLANG_ENABLE_LOGITS_PROCESSER_CHUNK = EnvBool(False)
SGLANG_ENABLE_LOGITS_PROCESSER_CHUNK = EnvBool(True)
SGLANG_LOGITS_PROCESSER_CHUNK_SIZE = EnvInt(2048)
# Tool-Call behavior
+15 -5
View File
@@ -793,9 +793,13 @@ class LogitsProcessor(nn.Module):
# 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
# 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)
chunk_logits = self._get_logits(
chunk_states, lm_head, logits_metadata, use_logits_buffer=False
)
# Initialize sampled_logits on first chunk
if i == 0:
@@ -896,6 +900,7 @@ class LogitsProcessor(nn.Module):
lm_head: VocabParallelEmbedding,
logits_metadata: LogitsMetadata,
embedding_bias: Optional[torch.Tensor] = None,
use_logits_buffer: bool = True,
) -> torch.Tensor:
"""Get logits from hidden_states.
@@ -922,7 +927,9 @@ class LogitsProcessor(nn.Module):
logits, local_hidden_states, logits_metadata
)
logits = self._copy_logits_to_buffer(logits, logits_metadata)
logits = self._copy_logits_to_buffer(
logits, logits_metadata, use_buffer=use_logits_buffer
)
if self.final_logit_softcapping:
if not (_is_npu or _is_cpu):
@@ -1038,9 +1045,12 @@ class LogitsProcessor(nn.Module):
return logits
def _copy_logits_to_buffer(
self, logits: torch.Tensor, logits_metadata: LogitsMetadata
self,
logits: torch.Tensor,
logits_metadata: LogitsMetadata,
use_buffer: bool = True,
) -> torch.Tensor:
logits_buffer = logits_metadata.next_token_logits_buffer
logits_buffer = logits_metadata.next_token_logits_buffer if use_buffer else None
if logits.shape[-1] > self.vocab_size:
logits = logits[:, : self.vocab_size]
logits_width = logits.shape[-1]
@@ -339,7 +339,7 @@ class TritonLoRABackend(BaseLoRABackend):
merged_segments = merge_and_chunk_segments(
seg_wi, seg_lens_list, chunk_size=pass_total
)
self.lm_head_pass_batch_infos.append(
lm_head_pass_batch_infos.append(
self._build_lm_head_batch_info(
merged_segments, batch_info, pass_total
)