Source: https://cs.stanford.edu/people/mmahoney/cs369m/Lectures/lecture1.pdf
(n: int, epsilon: float)
| 722 | samples = _fast_multinomial(probs, num_samples) |
| 723 | return samples |
| 724 | |
| 725 | def compute_logits( |
| 726 | self, |
| 727 | hidden_states: torch.Tensor, # [n_hidden_states, D] |
| 728 | ) -> torch.Tensor: |
| 729 | h_p = hidden_states @ self.rand_mat # [n_hidden_states, k] |
| 730 | return h_p @ self.w_p.T # [n_hidden_states, V] |