IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /max/get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /max/get-started.md).

Mojo function

topk_sampling_from_prob

def topk_sampling_from_prob[dtype: DType, out_idx_type: DType, block_size: Int = Int(1024), TopKArrLayoutType: TensorLayout = Layout[*?, *?], IndicesLayoutType: TensorLayout = Layout[*?, *?]](ctx: DeviceContext, probs: TileTensor[dtype, Storage=probs.Storage, address_space=probs.address_space, linear_idx_type=probs.linear_idx_type], output: TileTensor[out_idx_type, Storage=output.Storage, address_space=output.address_space, linear_idx_type=output.linear_idx_type], top_k_val: Int, deterministic: Bool = False, rng_seed: UInt64 = UInt64(0), rng_offset: UInt64 = UInt64(0), indices: Optional[TileTensor[out_idx_type, IndicesLayoutType, MutUntrackedOrigin]] = None, top_k_arr: Optional[TileTensor[out_idx_type, TopKArrLayoutType, MutUntrackedOrigin]] = None)

Top-K sampling from probability distribution.

Performs stochastic sampling from a probability distribution, considering only the top-k most probable tokens. Uses rejection sampling with ternary search to efficiently find appropriate samples.

Parameters:

  • ​dtype (DType): Element type of the probs tensor.
  • ​out_idx_type (DType): Index type used for the sampled output indices.
  • ​block_size (Int): Number of threads per block (defaults to 1024).
  • ​TopKArrLayoutType (TensorLayout): Memory layout of the optional top_k_arr tensor.
  • ​IndicesLayoutType (TensorLayout): Memory layout of the optional indices tensor.

Args:

Raises:

Error: If tensor ranks or shapes are invalid.