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
kentries oflexsort(-prob, index)(NaN/negative probabilities sanitized to+0). Its layout is deterministic (a pure function of the input) but not sorted; entries beyondcountare undefined. Same dispatch conditions astop_k_top_p_sampling_from_probs(); requests the frozen kernels cannot serve raiseValueErrorwith thecake_sampling_route()reason.- Parameters:
probs (torch.Tensor) –
float32 [batch, vocab]probabilities, contiguous.top_k (Union[int, torch.Tensor]) – Number of candidates kept per row (
intorint32 [batch]),1 <= k <= 1024.top_k_max (Optional[int]) – Upper bound of
top_kwhen 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]