Reduce across V-tiles (dim=1) on this rank and adjust to global vocab indices.
(
maxs: torch.Tensor, # [num_samples, n_tiles, H]
maxs_idx: torch.Tensor, # [num_samples, n_tiles, H]
vocab_start_index: int,
)
| 342 | |
| 343 | @torch.compile(fullgraph=True) |
| 344 | def _local_reduce( |
| 345 | maxs: torch.Tensor, # [num_samples, n_tiles, H] |
| 346 | maxs_idx: torch.Tensor, # [num_samples, n_tiles, H] |
| 347 | vocab_start_index: int, |
| 348 | ) -> tuple[torch.Tensor, torch.Tensor]: |
| 349 | """Reduce across V-tiles (dim=1) on this rank and adjust to global vocab indices.""" |
| 350 | idxs = maxs.max(dim=1).indices # [num_samples, H] |
| 351 | samples = maxs_idx.gather(1, idxs.unsqueeze(1)).squeeze(1) # [num_samples, H] |
| 352 | max_values = maxs.gather(1, idxs.unsqueeze(1)).squeeze(1) # [num_samples, H] |
| 353 | samples += vocab_start_index |
| 354 | return samples.T.contiguous(), max_values.T.contiguous() # [H, num_samples] |
| 355 | |
| 356 | |
| 357 | def clip(low, high, x): |
no outgoing calls
no test coverage detected