Baseline: matmul for logits followed by argmax. Returns [n_hidden_states, 1].
(
weights: torch.Tensor, # [V, D]
hidden_states: torch.Tensor, # [n_hidden_states, D]
num_samples: int, # ignored (always returns 1 sample per row)
temperature: torch.Tensor, # ignored (greedy)
tp: "TPInfo" = TP1,
**_kwargs,
)
| 137 | ) |
| 138 | logits = apply_top_k_top_p_triton(logits, k, p) |
| 139 | return logits.softmax(dim=-1) |
| 140 | |
| 141 | |
| 142 | @nvtx.annotate() |
| 143 | def greedy_sample( |
| 144 | weights: torch.Tensor, # [V, D] |
| 145 | hidden_states: torch.Tensor, # [n_hidden_states, D] |
| 146 | num_samples: int, # ignored (always returns 1 sample per row) |
| 147 | temperature: torch.Tensor, # ignored (greedy) |
| 148 | tp: "TPInfo" = TP1, |
| 149 | **_kwargs, |
| 150 | ) -> torch.Tensor: |
| 151 | """Baseline: matmul for logits followed by argmax. Returns [n_hidden_states, 1].""" |
| 152 | logits = hidden_states @ weights.T # [n_hidden_states, V] |
| 153 | if tp.size > 1: |
| 154 | logits = _allgather_logits(logits) |
nothing calls this directly
no test coverage detected