tilelang.ascend.language.tile_schedule ====================================== .. py:module:: tilelang.ascend.language.tile_schedule .. autoapi-nested-parse:: 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 ------- .. autoapisummary:: tilelang.ascend.language.tile_schedule.AscendBaseTileScheduler tilelang.ascend.language.tile_schedule.AscendTileScheduler tilelang.ascend.language.tile_schedule.AscendBatchedTileScheduler tilelang.ascend.language.tile_schedule.AscendMGroupedTileScheduler tilelang.ascend.language.tile_schedule.AscendKGroupedTileScheduler Module Contents --------------- .. py:class:: AscendBaseTileScheduler(*, block_m, block_n, num_cores, shape_m, shape_n, group_m = 4, xor_n = True, snake = True, stateful = True, name = None) Bases: :py:obj:`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() :param block_m: Tile sizes along M and N. :type block_m: int :param block_n: Tile sizes along M and N. :type block_n: int :param num_cores: Persistent stride = number of resident cores (grid dim). :type num_cores: int | PrimExpr :param shape_m: Problem M and N. :type shape_m: int | PrimExpr :param shape_n: Problem M and N. :type shape_n: int | PrimExpr :param group_m: Swizzle panel height along M (4 or 8). :type group_m: int :param xor_n: XOR the in-panel n-index with the m-index (default True). :type xor_n: bool :param snake: Reverse the N-group order on odd superrows (default True). :type snake: bool :param stateful: If True (default) allocate state and enable ``init`` / ``valid`` / ``next_block``. If False expose only the store-free ``coord`` decode. :type stateful: bool :param name: Optional state-buffer name prefix. The default ``None`` uses generic auto-generated names. :type name: str, optional .. rubric:: Notes State lives in single-element ``T.alloc_var`` buffers -- read with ``sched.m_idx[0]`` and (inside methods) write with ``self.x[0] = ...``. .. py:attribute:: core_id :value: None .. py:attribute:: valid_flag :value: None .. py:method:: 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. .. py:method:: coord(block_idx) Decode a linear ``block_idx`` into ``(m_block, n_block)`` (Normal grid). .. py:method:: update_current_idx(linear_idx) .. py:method:: valid() .. py:method:: init(core_id) .. py:method:: next_block() .. py:method:: get_actual_m(m_block_idx) .. py:method:: get_actual_n(n_block_idx) .. py:class:: AscendTileScheduler(*, block_m, block_n, num_cores, shape_m, shape_n, group_m = 4, xor_n = True, snake = True, stateful = True, name = None) Bases: :py:obj:`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. .. py:class:: 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: :py:obj:`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. .. py:method:: get_batch_idx() .. py:method:: 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. .. py:class:: 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: :py:obj:`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). .. py:method:: get_group_idx() .. py:method:: get_actual_m(m_block_idx) .. py:method:: 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 + ``, so it stays non-negative and every full tile fits the padded row space. .. py:class:: 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: :py:obj:`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). .. py:method:: get_group_idx() .. py:method:: get_k_idx_base() .. py:method:: get_sf_idx_base() .. py:method:: get_shape_k() .. py:method:: 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.