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
GroupedTileScheduler
struct GroupedTileScheduler[tile_m: Int, tile_n: Int, tile_k: Int, max_groups: Int, num_stages: Int = Int(0), cta_group: Int = Int(1), *, problem_sizes_engine: TensorEngine]
Tile scheduler for grouped block-scaled GEMM.
Uses linear tile iteration to map tiles across groups. Does not use CLC (Cluster Launch Control) since work distribution is deterministic.
Parameters
- tile_m (
Int): M dimension of output tiles. - tile_n (
Int): N dimension of output tiles. - tile_k (
Int): K dimension of input tiles. - max_groups (
Int): Maximum number of groups. - num_stages (
Int): Pipeline stages (0 = single wave). - cta_group (
Int): Number of CTAs cooperating per tile (1 or 2 for 2SM). - problem_sizes_engine (
TensorEngine): Engine of the problem-sizes tile.
Fields
- num_groups (
Int): Number of active groups. - problem_sizes (
TileTensor[.int32, Layout[TypeList[ComptimeInt[max_groups], ComptimeInt[Int(4)]](), TypeList[ComptimeInt[Int(4)], ComptimeInt[Int(1)]]()], MutAnyOrigin, Engine=problem_sizes_engine]): Problem sizes tensor (num_groups, 4) with [M, N, K, L] per group.
Implemented traits
AnyType,
Copyable,
Deinitable,
ImplicitlyCopyable,
Movable,
RegisterPassable,
TrivialRegisterPassable
Methods
__init__
def __init__(problem_sizes: TileTensor[.int32, Layout[TypeList[ComptimeInt[max_groups], ComptimeInt[Int(4)]](), TypeList[ComptimeInt[Int(4)], ComptimeInt[Int(1)]]()], MutAnyOrigin, Engine=problem_sizes_engine], num_groups: Int) -> Self
Initialize scheduler with problem sizes.
Args:
- problem_sizes (
TileTensor[.int32, Layout[TypeList[ComptimeInt[max_groups], ComptimeInt[Int(4)]](), TypeList[ComptimeInt[Int(4)], ComptimeInt[Int(1)]]()], MutAnyOrigin, Engine=problem_sizes_engine]): (num_groups, 4) tensor with [M, N, K, L] per group. - num_groups (
Int): Number of active groups.
work_iterator
def work_iterator(self) -> GroupedWorkIterator[tile_m, tile_n, tile_k, max_groups, cta_group, problem_sizes_engine=problem_sizes_engine]
Create a per-warp work iterator.
Each warp should create its own work iterator. The iterator owns work_info and cumulative tile counts internally.
For 2SM (cta_group=2), the iterator uses cluster-based indexing.
Returns:
GroupedWorkIterator[tile_m, tile_n, tile_k, max_groups, cta_group, problem_sizes_engine=problem_sizes_engine]