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¶

AscendBaseTileScheduler

Base for persistent GEMM block schedulers (GROUP_M swizzle + N-first snake).

AscendTileScheduler

Single flat M x N grid, GROUP_M swizzle + N-first snake; the last

AscendBatchedTileScheduler

Standard batched matmul: the Normal grid replicated num_groups (= batch)

AscendMGroupedTileScheduler

MoE m-grouped "contiguous" layout. grouped_layout is the prefix-sum-of-rows

AscendKGroupedTileScheduler

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.BaseTileScheduler

Base for persistent GEMM block schedulers (GROUP_M swizzle + N-first snake).

Holds everything shared across GEMM variants: the grid geometry, the stateless _swizzle / coord decode, the persistent-block-index advance, and the init / valid / next_block loop protocol. Concrete variants subclass this and implement _step (advance one block + set m_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 iteration current_iter is current_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-free coord decode.

  • name (str, optional) – Optional state-buffer name prefix. The default None uses generic auto-generated names.

Notes

State lives in single-element T.alloc_var buffers – read with sched.m_idx[0] and (inside methods) write with self.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_n follow get_actual_m / get_actual_n, and every tile _step hands out is a grid tile with a non-empty remainder, so the emitted facts hold for every iteration of the while 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_idx into (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: AscendBaseTileScheduler

Single flat M x N grid, GROUP_M swizzle + N-first snake; the last < group_m M-rows fall back to row-major. This is the plain (non-grouped, non-batched) GEMM scheduler.

Parameters:
  • group_m (int)

  • xor_n (bool)

  • snake (bool)

  • stateful (bool)

  • name (str | None)

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: AscendBaseTileScheduler

Standard 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:
  • num_groups (int)

  • group_m (int)

  • xor_n (bool)

  • snake (bool)

  • stateful (bool)

  • name (str | None)

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: AscendBaseTileScheduler

MoE m-grouped “contiguous” layout. grouped_layout is the prefix-sum-of-rows array (length num_groups); group g’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 (GM int32 prefix-sum buffer), num_groups, alignment (group-start row alignment, default 256).

Parameters:
  • num_groups (int)

  • alignment (int)

  • group_m (int)

  • xor_n (bool)

  • snake (bool)

  • stateful (bool)

  • name (str | None)

get_group_idx()¶
get_actual_m(m_block_idx)¶
group()¶

Bind the current group index and assume its range.

valid gates the loop body on group_idx < num_groups; the m cursor is last_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: AscendBaseTileScheduler

MoE k-grouped. All groups share one M x N grid but differ in K length; grouped_layout[i] is the cumulative end-K of group i. 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 (GM int32 prefix-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:
  • num_groups (int)

  • alignment (int)

  • group_m (int)

  • xor_n (bool)

  • snake (bool)

  • stateful (bool)

  • name (str | None)

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 an alignment boundary, which is a multiple of the SF divisor.