(logits, top_p)
| 27 | |
| 28 | |
| 29 | def apply_top_p(logits, top_p): |
| 30 | probs = F.softmax(logits, dim=-1) |
| 31 | sorted_probs, sorted_indices = torch.sort(probs, descending=True, dim=-1) |
| 32 | cumulative_probs = torch.cumsum(sorted_probs, dim=-1) |
| 33 | sorted_indices_to_remove = cumulative_probs > top_p |
| 34 | sorted_indices_to_remove[..., 1:] = sorted_indices_to_remove[..., :-1].clone() |
| 35 | sorted_indices_to_remove[..., 0] = False |
| 36 | batch_size = logits.shape[0] |
| 37 | filtered_logits = logits.clone() |
| 38 | for i in range(batch_size): |
| 39 | indices_to_remove = sorted_indices[i][sorted_indices_to_remove[i]] |
| 40 | filtered_logits[i, indices_to_remove] = float("-inf") |
| 41 | return filtered_logits |
| 42 | |
| 43 | |
| 44 | def apply_top_p_optimized(logits, top_p): |
nothing calls this directly
no outgoing calls
no test coverage detected