flashinfer.cake_sampling.top_k_top_p_sampling_from_probs¶
- flashinfer.cake_sampling.top_k_top_p_sampling_from_probs(probs: Tensor, top_k: int | Tensor, top_p: float | Tensor, *, top_k_max: int | None = None, generator: Generator | None = None, philox_seed: int | None = None, philox_offset: int | None = None, out: Tensor | None = None, renorm_out: Tensor | None = None, workspace: tuple[Tensor, Tensor, Tensor] | None = None, enable_pdl: bool = True) Tensor¶
Fused top-k-then-top-p sampling from probabilities (thread-block-cluster radix pipeline).
- 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_p (Union[float, torch.Tensor]) – Nucleus threshold in
(0, 1](floatorfloat32 [batch]).top_k_max (Optional[int]) – Upper bound of
top_kwhen it is a tensor (avoids a device synchronization).generator (Optional[torch.Generator]) – Source of the Philox seed/offset (default CUDA generator when omitted), advanced exactly like
flashinfer.sampling.top_k_top_p_sampling_from_probs().philox_seed (Optional[int]) – Explicit Philox parameters (both required together);
generatoris then not touched.philox_offset (Optional[int]) – Explicit Philox parameters (both required together);
generatoris then not touched.out (Optional[torch.Tensor]) –
int32 [batch]output buffer (allocated when omitted).renorm_out (Optional[torch.Tensor]) – Optional
float32 [batch, 1024]buffer that receives the renormalized kept probabilities in sorted slab order (descending probability, ascending index; zeros for dropped entries; only the firstcountentries of a row are written). Served by the pipeline route only.workspace (Optional[tuple[torch.Tensor, torch.Tensor, torch.Tensor]]) – Stage-1 slab buffers
(values float32 [batch, 1024], indices int32 [batch, 1024], counts int32 [batch]). When given, the sorted top-k slab of this call is left in them (aligned withrenorm_out); otherwise a cached per-shape workspace is used.enable_pdl (bool) – Launch stage 2/3 with programmatic dependent launch.
- Returns:
samples –
int32 [batch]sampled token ids.- Return type:
torch.Tensor