(probs, p)
| 86 | logprobs: List[float] # not required |
| 87 | |
| 88 | def sample_top_p(probs, p): |
| 89 | probs_sort, probs_idx = torch.sort(probs, dim=-1, descending=True) |
| 90 | probs_sum = torch.cumsum(probs_sort, dim=-1) |
| 91 | mask = probs_sum - probs_sort > p |
| 92 | probs_sort[mask] = 0.0 |
| 93 | probs_sort.div_(probs_sort.sum(dim=-1, keepdim=True)) |
| 94 | next_token = torch.multinomial(probs_sort, num_samples=1) |
| 95 | next_token = torch.gather(probs_idx, -1, next_token) |
| 96 | return next_token |
| 97 | |
| 98 | |
| 99 | Dialog = List[Message] |
nothing calls this directly
no outgoing calls
no test coverage detected