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
|
self, logits: torch.Tensor, logits_metadata: LogitsMetadata
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
logits_buffer = logits_metadata.next_token_logits_buffer
|
logits_buffer = logits_metadata.next_token_logits_buffer
|
||||||
# The shared logits buffer is keyed by vocab width; skip it when this
|
if logits.shape[-1] > self.vocab_size:
|
||||||
# model's vocab doesn't match (e.g. hot-vocab draft vs full-vocab target).
|
logits = logits[:, : self.vocab_size]
|
||||||
if logits_buffer is not None and logits_buffer.shape[-1] == 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
|
assert logits_buffer.dtype == torch.float
|
||||||
logits_buffer.copy_(logits[:, : self.vocab_size])
|
logits_buffer.copy_(logits)
|
||||||
logits = logits_buffer
|
logits = logits_buffer
|
||||||
else:
|
else:
|
||||||
logits = logits[:, : self.vocab_size].float()
|
logits = logits.float()
|
||||||
return logits
|
return logits
|
||||||
|
|
||||||
def _get_dllm_logits(
|
def _get_dllm_logits(
|
||||||
|
|||||||
Reference in New Issue
Block a user