24 lines
947 B
Python
24 lines
947 B
Python
from __future__ import annotations
|
|
|
|
import torch
|
|
|
|
|
|
@torch.inference_mode()
|
|
def sample_next_token(logits: torch.Tensor, temperature: float = 0.7, top_p: float = 0.95) -> int:
|
|
logits = logits[0, -1, :].float()
|
|
if temperature is None or temperature <= 0:
|
|
return int(torch.argmax(logits).item())
|
|
logits = logits / float(temperature)
|
|
probs = torch.softmax(logits, dim=-1)
|
|
if top_p is not None and 0 < top_p < 1:
|
|
sorted_probs, sorted_indices = torch.sort(probs, descending=True)
|
|
cumulative = torch.cumsum(sorted_probs, dim=-1)
|
|
mask = cumulative > top_p
|
|
mask[1:] = mask[:-1].clone()
|
|
mask[0] = False
|
|
sorted_probs = sorted_probs.masked_fill(mask, 0.0)
|
|
sorted_probs = sorted_probs / sorted_probs.sum()
|
|
idx = torch.multinomial(sorted_probs, num_samples=1)
|
|
return int(sorted_indices[idx].item())
|
|
return int(torch.multinomial(probs, num_samples=1).item())
|