from __future__ import annotations import torch def sample_token( logits: torch.Tensor, *, do_sample: bool = False, top_k: int = 50, top_p: float = 1.0, temperature: float = 0.9, generator: torch.Generator | None = None, ) -> torch.Tensor: if not do_sample: return logits.argmax(dim=-1) scores = logits.float() / temperature if 0 < top_k < scores.shape[-1]: threshold = scores.topk(top_k, dim=-1).values[..., -1, None] scores = scores.masked_fill(scores < threshold, -torch.inf) if top_p < 1.0: sorted_scores, sorted_indices = scores.sort(dim=-1, descending=True) remove = sorted_scores.softmax(dim=-1).cumsum(dim=-1) > top_p remove[..., 1:] = remove[..., :-1].clone() remove[..., 0] = False sorted_scores = sorted_scores.masked_fill(remove, -torch.inf) scores = torch.full_like(scores, -torch.inf).scatter( -1, sorted_indices, sorted_scores ) return torch.multinomial( scores.softmax(dim=-1), num_samples=1, generator=generator ).squeeze(-1)