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 (int or int32 [batch]), 1 <= k <= 1024.

  • top_p (Union[float, torch.Tensor]) – Nucleus threshold in (0, 1] (float or float32 [batch]).

  • top_k_max (Optional[int]) – Upper bound of top_k when 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); generator is then not touched.

  • philox_offset (Optional[int]) – Explicit Philox parameters (both required together); generator is 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 first count entries 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 with renorm_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