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 struct
TopKPerRow
struct TopKPerRow
Registers the mo.top_k.per_row graph op with the graph compiler.
mo.top_k over the last axis of a [rows, n] input, with a count per
row: row r keeps the first k[r] picks mo.top_k with max_k makes
for it, in the same order, and pads the rest of its max_k slots with
the dead value and index -1. The kernel stops after k[r] picks, so
its cost follows the counts rather than max_k.
Implemented traits
Methods
execute
static def execute[dtype: DType, //, max_k: Int, target: StringSpan[ImmStaticOrigin]](values: ManagedTensorSlice[IOSpec[_, _].Output, static_spec=values.static_spec], indices: ManagedTensorSlice[IOSpec[_, _].Output, static_spec=indices.static_spec], input: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=input.static_spec], k: ManagedTensorSlice[IOSpec[_, _].Input, static_spec=k.static_spec], ctx: DeviceContext)
Executes the mo.top_k.per_row graph op.
Parameters:
- dtype (
DType): Element type of the input and values. - max_k (
Int): Output width, the largest count any row may ask for. - target (
StringSpan[ImmStaticOrigin]): Compilation target string.
Args:
- values (
ManagedTensorSlice[IOSpec[_, _].Output, static_spec=values.static_spec]):[rows, max_k]selected values. - indices (
ManagedTensorSlice[IOSpec[_, _].Output, static_spec=indices.static_spec]):[rows, max_k]their column indices. - input (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=input.static_spec]):[rows, n]values to select from. - k (
ManagedTensorSlice[IOSpec[_, _].Input, static_spec=k.static_spec]):[rows]picks per row, each in[0, max_k]. - ctx (
DeviceContext): Device context used to enqueue the kernel.
Raises:
Error: If the operation parameters are invalid.