For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).
Mojo function
TopKTopPSamplingEmitDistClusterKernel
def TopKTopPSamplingEmitDistClusterKernel[ProbsLayoutType: TensorLayout, probs_origin: ImmOrigin, OutputLayoutType: TensorLayout, output_origin: MutOrigin, block_size: Int, vec_size: Int, dtype: DType, out_idx_type: DType, deterministic: Bool, cluster_size: Int, dist_dtype: DType = .float32, ProbsEngine: TensorEngine = DefaultEngine, OutputEngine: TensorEngine = DefaultEngine](probs: TileTensor[dtype, ProbsLayoutType, probs_origin, Engine=ProbsEngine], output: TileTensor[out_idx_type, OutputLayoutType, output_origin, Engine=OutputEngine], out_dist: Pointer[Scalar[dist_dtype], MutAnyOrigin], indices: Optional[Pointer[Scalar[out_idx_type], ImmutAnyOrigin]], top_k_arr: Optional[Pointer[Scalar[out_idx_type], ImmutAnyOrigin]], top_k_val: Int32, top_p_arr: Optional[Pointer[Float32, ImmutAnyOrigin]], top_p_val: Float32, d: Int32, rng_seed: Optional[Pointer[UInt64, ImmutAnyOrigin]], rng_offset: UInt64, temperature: Optional[Pointer[Float32, ImmutAnyOrigin]], min_p: Optional[Pointer[Float32, ImmutAnyOrigin]])
TopKTopPSamplingFromProbKernel with from_logits and emit_dist, spread over a cluster.
Every phase shares the row across the cluster. Each CTA stages its
min-p-masked slice of the softmax weights in shared memory once, and the
rejection loop, the cutoff search and the mask all read that copy: the
loop's per-trial CDF prefix decomposes over the contiguous slices (see
_sampling_rejection_loop_cluster), and its pivot masses combine across
the cluster exactly like the search's statistics.
Everything downstream of the loop consumes the loop's own combined
values -- row_max, p_eff, the bracket and its mass budget -- on every
CTA. Rebuilding any of them independently risks a one-ulp disagreement
that puts the sampled token outside the emitted nucleus, which is the
hazard the fused single-block kernel exists to avoid.
The launch must set cluster_dim to cluster_size.
Parameters:
- ProbsLayoutType (
TensorLayout): Memory layout of the inputprobstile. - probs_origin (
ImmOrigin): Origin tag for the immutable inputprobstile. - OutputLayoutType (
TensorLayout): Memory layout of the outputoutputtile. - output_origin (
MutOrigin): Origin tag for the mutable outputoutputtile. - block_size (
Int): Number of threads per block. - vec_size (
Int): Number of elements each thread loads per vectorized access. - dtype (
DType): Element type of theprobstensor. - out_idx_type (
DType): Index type used for the sampled output indices. - deterministic (
Bool): If True, use deterministic sampling. - cluster_size (
Int): Number of CTAs sharing each row. - dist_dtype (
DType): Element type ofout_dist. - ProbsEngine (
TensorEngine): Engine of the inputprobstile. - OutputEngine (
TensorEngine): Engine of the outputoutputtile.
Args:
- probs (
TileTensor[dtype, ProbsLayoutType, probs_origin, Engine=ProbsEngine]): Input logits [batch_size, d]. - output (
TileTensor[out_idx_type, OutputLayoutType, output_origin, Engine=OutputEngine]): Output sampled indices [batch_size]. - out_dist (
Pointer[Scalar[dist_dtype], MutAnyOrigin]): Output masked renormalized distribution [batch_size, d]. - indices (
Optional[Pointer[Scalar[out_idx_type], ImmutAnyOrigin]]): Optional row indices for batch indexing [batch_size]. - top_k_arr (
Optional[Pointer[Scalar[out_idx_type], ImmutAnyOrigin]]): Optional per-row top_k values [batch_size]. - top_k_val (
Int32): Default top_k value if top_k_arr is null. - top_p_arr (
Optional[Pointer[Float32, ImmutAnyOrigin]]): Optional per-row top_p values [batch_size]. - top_p_val (
Float32): Default top_p value if top_p_arr is null. - d (
Int32): Vocabulary size. - rng_seed (
Optional[Pointer[UInt64, ImmutAnyOrigin]]): Optional per-row seed array [batch_size], indexed by row_idx. If null, defaults to 0. - rng_offset (
UInt64): Random offset for Random number generator. - temperature (
Optional[Pointer[Float32, ImmutAnyOrigin]]): Optional per-row temperature [batch_size]; 0 is clamped. - min_p (
Optional[Pointer[Float32, ImmutAnyOrigin]]): Optional per-row min-p thresholds [batch_size].