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
TopKTopPMaskedProbsKernel
def TopKTopPMaskedProbsKernel[block_size: Int, vec_size: Int, dtype: DType, LogitsLayoutType: TensorLayout, logits_origin: ImmOrigin, coop_size: Int = Int(1)](logits: TileTensor[dtype, LogitsLayoutType, logits_origin], probs_ptr: Pointer[Float32, MutAnyOrigin], top_k_arr: Optional[Pointer[Int64, ImmutAnyOrigin]], top_k_val: Int32, top_p_arr: Optional[Pointer[Float32, ImmutAnyOrigin]], top_p_val: Float32, temperature: Optional[Pointer[Float32, ImmutAnyOrigin]], d: Int32, coop_ws: Optional[Pointer[Int32, MutAnyOrigin]])
Writes each row's top-k/top-p masked softmax.
With coop_size above one, coop_size blocks share a row: each owns a
contiguous slice and combines whole-row statistics through coop_ws. The
launch must keep every block in a group resident.
Parameters:
- block_size (
Int): Number of threads per block. - vec_size (
Int): Number of elements each thread loads per access. - dtype (
DType): Element type oflogits. - LogitsLayoutType (
TensorLayout): Memory layout of thelogitstile. - logits_origin (
ImmOrigin): Origin tag for the immutablelogitstile. - coop_size (
Int): Blocks sharing each row; 1 keeps the whole row in one block and compiles the cross-block traffic away.
Args:
- logits (
TileTensor[dtype, LogitsLayoutType, logits_origin]): Input logits [batch_size, d]. - probs_ptr (
Pointer[Float32, MutAnyOrigin]): Output masked renormalized distribution [batch_size, d]. - top_k_arr (
Optional[Pointer[Int64, ImmutAnyOrigin]]): Optional per-row top-k [batch_size]. - top_k_val (
Int32): Default top-k iftop_k_arris null. - top_p_arr (
Optional[Pointer[Float32, ImmutAnyOrigin]]): Optional per-row top-p [batch_size]. - top_p_val (
Float32): Default top-p iftop_p_arris null. - temperature (
Optional[Pointer[Float32, ImmutAnyOrigin]]): Optional per-row temperature [batch_size]. - d (
Int32): Vocabulary size. - coop_ws (
Optional[Pointer[Int32, MutAnyOrigin]]): Zeroed cross-block workspace; required whencoop_sizeexceeds one, ignored otherwise.