tilelang.ascend.language.tile_schedule¶
Ascend persistent GEMM block schedulers.
Built on the shared BaseTileScheduler skeleton in
tilelang.language.tile_schedule; see that module’s docstring for how
@meta_class inlining and the T.alloc_var state protocol work.
Next to the raw state accessors, a scheduler hands out the tile scalars a kernel
slices with: tile(), plus batch() / group() on the variants. Each
call binds its values, assumes their range contract, and returns the bound
variables:
scheduler.init(block_idx)
while scheduler.valid():
m_idx, n_idx, actual_m, actual_n = scheduler.tile()
...
Those assumptions are what lets the pre-schedule passes prove per-tile GM regions and L1/L0 extents in bounds.
Classes¶
Base for persistent GEMM block schedulers (GROUP_M swizzle + N-first snake). |
|
Single flat M x N grid, GROUP_M swizzle + N-first snake; the last |
|
Standard batched matmul: the Normal grid replicated |
|
MoE m-grouped "contiguous" layout. |
|
MoE k-grouped. All groups share one M x N grid but differ in K length; |
Module Contents¶
- class tilelang.ascend.language.tile_schedule.AscendBaseTileScheduler(*, block_m, block_n, num_cores, shape_m, shape_n, group_m=4, xor_n=True, snake=True, stateful=True, name=None)¶
Bases:
tilelang.language.tile_schedule.BaseTileSchedulerBase for persistent GEMM block schedulers (GROUP_M swizzle + N-first snake).
Holds everything shared across GEMM variants: the grid geometry, the stateless
_swizzle/coorddecode, the persistent-block-index advance, and theinit/valid/next_blockloop protocol. Concrete variants subclass this and implement_step(advance one block + setm_idx/n_idx/valid_flag); grouped variants add their own state and_init_state.Each persistent core independently walks the global
(m_block, n_block)grid. The block index handed out on iterationcurrent_iteriscurrent_iter * num_cores + core_id, mapped to a tile by_swizzle. Usage:sched = T.AscendTileScheduler(block_m=BM, block_n=BN, num_cores=NUM_CORES, shape_m=M, shape_n=N) sched.init(bx) while sched.valid(): m_tile, n_tile = sched.m_idx[0], sched.n_idx[0] # ... compute tile (m_tile, n_tile) ... sched.next_block()
- Parameters:
block_m (int) – Tile sizes along M and N.
block_n (int) – Tile sizes along M and N.
num_cores (int | PrimExpr) – Persistent stride = number of resident cores (grid dim).
shape_m (int | PrimExpr) – Problem M and N.
shape_n (int | PrimExpr) – Problem M and N.
group_m (int) – Swizzle panel height along M (4 or 8).
xor_n (bool) – XOR the in-panel n-index with the m-index (default True).
snake (bool) – Reverse the N-group order on odd superrows (default True).
stateful (bool) – If True (default) allocate state and enable
init/valid/next_block. If False expose only the store-freecoorddecode.name (str, optional) – Optional state-buffer name prefix. The default
Noneuses generic auto-generated names.
Notes
State lives in single-element
T.alloc_varbuffers – read withsched.m_idx[0]and (inside methods) write withself.x[0] = ....- core_id = None¶
- valid_flag = None¶
- tile()¶
Bind the current tile’s M/N offsets and extents, and assume their ranges.
Returns
(m_idx, n_idx, actual_m, actual_n): the tile’s element offsets along M/N and the in-bounds extents of its M/N tail.actual_m/actual_nfollowget_actual_m/get_actual_n, and every tile_stephands out is a grid tile with a non-empty remainder, so the emitted facts hold for every iteration of thewhile sched.valid()loop they dominate. Call it once at the top of that loop body and slice with the returned values.
- coord(block_idx)¶
Decode a linear
block_idxinto(m_block, n_block)(Normal grid).
- update_current_idx(linear_idx)¶
- valid()¶
- init(core_id)¶
- next_block()¶
- get_actual_m(m_block_idx)¶
- get_actual_n(n_block_idx)¶
- class tilelang.ascend.language.tile_schedule.AscendTileScheduler(*, block_m, block_n, num_cores, shape_m, shape_n, group_m=4, xor_n=True, snake=True, stateful=True, name=None)¶
Bases:
AscendBaseTileSchedulerSingle flat M x N grid, GROUP_M swizzle + N-first snake; the last
< group_mM-rows fall back to row-major. This is the plain (non-grouped, non-batched) GEMM scheduler.
- class tilelang.ascend.language.tile_schedule.AscendBatchedTileScheduler(*, block_m, block_n, num_cores, shape_m, shape_n, num_groups, group_m=4, xor_n=True, snake=True, stateful=True, name=None)¶
Bases:
AscendBaseTileSchedulerStandard batched matmul: the Normal grid replicated
num_groups(= batch) times.get_batch_idx()exposes the batch of the current block.Extra parameter
num_groups(= batch count) beyond the base scheduler.- Parameters:
- get_batch_idx()¶
- batch()¶
Bind the current batch index and assume its range.
Returns the same bound variable the fact is stated over; use it (instead of the raw
get_batch_idx()read) to slice per-batch GM regions.
- class tilelang.ascend.language.tile_schedule.AscendMGroupedTileScheduler(*, block_m, block_n, num_cores, shape_m, shape_n, grouped_layout, num_groups, alignment=256, group_m=4, xor_n=True, snake=True, stateful=True, name=None)¶
Bases:
AscendBaseTileSchedulerMoE m-grouped “contiguous” layout.
grouped_layoutis the prefix-sum-of-rows array (lengthnum_groups); groupg’s valid rows are[align(psum[g-1], alignment), psum[g])in the global M space (group 0 starts at 0).get_group_idx()exposes the current group for per-group B/SF offsets. A/D must allocate all aligned rows, and physical M must be a multiple of alignment.Extra parameters beyond the base:
grouped_layout(GMint32prefix-sum buffer),num_groups,alignment(group-start row alignment, default 256).- Parameters:
- get_group_idx()¶
- get_actual_m(m_block_idx)¶
- group()¶
Bind the current group index and assume its range.
validgates the loop body ongroup_idx < num_groups; the m cursor islast_psum_m // block_m + <group-local tile>, so it stays non-negative and every full tile fits the padded row space.
- class tilelang.ascend.language.tile_schedule.AscendKGroupedTileScheduler(*, block_m, block_n, num_cores, shape_m, shape_n, shape_k, grouped_layout, num_groups, alignment=256, group_m=4, xor_n=True, snake=True, stateful=True, name=None)¶
Bases:
AscendBaseTileSchedulerMoE k-grouped. All groups share one M x N grid but differ in K length;
grouped_layout[i]is the cumulative end-K of groupi.get_group_idx(),get_k_idx_base()/get_sf_idx_base()give the group and its A/B and SF K offsets;get_shape_k()is the current group’s K length.Extra parameters beyond the base:
grouped_layout(GMint32prefix-sum buffer),num_groups,shape_k(physical K allocation, including trailing padding; only the K-range contract uses it),alignment(group-start K alignment, default 256).- Parameters:
- get_group_idx()¶
- get_k_idx_base()¶
- get_sf_idx_base()¶
- get_shape_k()¶
- group()¶
Bind the current group’s K window and assume its ranges.
Returns
(group_idx, k_base, group_k, sf_base): the group index, the group’s K origin, its K length, and the origin of its packed scale-factor rows. A valid tile belongs to a non-empty group, so the K span[k_base, k_base + group_k)stays inside the total K and its SF rows stay inside the packed SF buffer: a group starts on analignmentboundary, which is a multiple of the SF divisor.