diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index e318e908f..142e90543 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -981,14 +981,19 @@ class LogitsProcessor(nn.Module): self, logits: torch.Tensor, logits_metadata: LogitsMetadata ) -> torch.Tensor: logits_buffer = logits_metadata.next_token_logits_buffer - # The shared logits buffer is keyed by vocab width; skip it when this - # model's vocab doesn't match (e.g. hot-vocab draft vs full-vocab target). - if logits_buffer is not None and logits_buffer.shape[-1] == self.vocab_size: + if logits.shape[-1] > self.vocab_size: + logits = logits[:, : self.vocab_size] + logits_width = logits.shape[-1] + # The shared logits buffer is keyed by vocab width and rows; skip it + # when this batch has a different logits shape than the graph buffer. + if logits_buffer is not None and tuple(logits_buffer.shape) == tuple( + logits.shape + ): assert logits_buffer.dtype == torch.float - logits_buffer.copy_(logits[:, : self.vocab_size]) + logits_buffer.copy_(logits) logits = logits_buffer else: - logits = logits[:, : self.vocab_size].float() + logits = logits.float() return logits def _get_dllm_logits(