Fix shared logits buffer for reduced-vocab draft models (#29943)

This commit is contained in:
Po-Han Huang (NVIDIA)
2026-07-02 16:08:56 -07:00
committed by GitHub
parent 8519be82e8
commit 17cce6a85f
+10 -5
View File
@@ -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(