(logits, thres=0.9)
| 25 | |
| 26 | |
| 27 | def top_k(logits, thres=0.9): |
| 28 | k = int((1 - thres) * logits.shape[-1]) |
| 29 | val, ind = torch.topk(logits, k) |
| 30 | probs = torch.full_like(logits, float("-inf")) |
| 31 | probs.scatter_(1, ind, val) |
| 32 | return probs |
| 33 | |
| 34 | |
| 35 | class AutoregressiveWrapper(nn.Module): |