Fix shared logits buffer for reduced-vocab draft models (#29943)
This commit is contained in:
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user