File size: 1,098 Bytes
d81a465
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
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)