flashinfer.cake_sampling.top_k_probs_to_slab

flashinfer.cake_sampling.top_k_probs_to_slab(probs: Tensor, top_k: int | Tensor, *, top_k_max: int | None = None, out_vals: Tensor | None = None, out_idx: Tensor | None = None, out_count: Tensor | None = None) → tuple[Tensor, Tensor, Tensor]

Stage 1 alone: exact per-row top-k into a [batch, 1024] slab.

The slab holds the first k entries of lexsort(-prob, index) (NaN/negative probabilities sanitized to +0). Its layout is deterministic (a pure function of the input) but not sorted; entries beyond count are undefined. Same dispatch conditions as top_k_top_p_sampling_from_probs(); requests the frozen kernels cannot serve raise ValueError with the cake_sampling_route() reason.

Parameters:
  • probs (torch.Tensor) – float32 [batch, vocab] probabilities, contiguous.

  • top_k (Union[int, torch.Tensor]) – Number of candidates kept per row (int or int32 [batch]), 1 <= k <= 1024.

  • top_k_max (Optional[int]) – Upper bound of top_k when it is a tensor (avoids a device synchronization).

  • out_vals (Optional[torch.Tensor]) – float32 [batch, 1024] slab values buffer (allocated when omitted).

  • out_idx (Optional[torch.Tensor]) – int32 [batch, 1024] slab indices buffer (allocated when omitted).

  • out_count (Optional[torch.Tensor]) – int32 [batch] per-row kept counts buffer (allocated when omitted).

Returns:

values, indices, counts – The slab (values [batch, 1024], indices [batch, 1024], counts [batch]).

Return type:

tuple[torch.Tensor, torch.Tensor, torch.Tensor]