tilelang.cuda.intrinsics.macro.wgmma_macro_generator¶

Attributes¶

Classes¶

WGMMADescriptorParams

Pre-computed WGMMA descriptor parameters, produced by

TensorCoreIntrinEmitter

To eliminate Python syntax within TIR Macro.

Functions¶

decode_k_panel_elems(k_mode, k_major, swizzle_mode, ...)

Element stride between consecutive K swizzle-atom panels, off the layout.

select_wgmma_inst_n(warp_col_tiles)

Widest legal WGMMA N that tiles warp_col_tiles exactly.

compute_gmma_descriptor(tl_layout, buffer, transposed)

Port of CuTe make_gmma_desc.

Module Contents¶

tilelang.cuda.intrinsics.macro.wgmma_macro_generator.lift¶
tilelang.cuda.intrinsics.macro.wgmma_macro_generator.decode_k_panel_elems(k_mode, k_major, swizzle_mode, swizzle_atom_elems, kind)¶

Element stride between consecutive K swizzle-atom panels, off the layout.

Shared by the WGMMA and UMMA descriptor builders. A K-major operand wider than one swizzle atom decodes with its K mode split into (atom, panels), and that panel sub-mode’s stride is the step the atom offset formulas need. It cannot be reconstructed as mn_extent * swizzle_atom_elems, which only holds when the operand covers its whole buffer – for a slice the MN extent shrinks while the panel spacing does not.

k_mode is the operand tile’s K mode in element space (post-restrict). Returns None only when no panel step is needed: an MN-major operand, or a layout with a single K panel, where the ki // k_atom_size factor the step is multiplied by is always zero.

A swizzled multi-panel layout whose K mode is not the canonical (atom, panels):(1, panel_stride) asserts rather than falling back to the extent-based reconstruction: that fallback is precisely the defect this decode exists to remove, so re-entering it silently would reintroduce a wrong-code path.

Parameters:
  • k_major (bool)

  • swizzle_atom_elems (int)

  • kind (str)

Return type:

int | None

tilelang.cuda.intrinsics.macro.wgmma_macro_generator.select_wgmma_inst_n(warp_col_tiles)¶

Widest legal WGMMA N that tiles warp_col_tiles exactly.

Hopper WGMMA accepts N in [8, 256] with N % 8 == 0, so an extent such as 96 or 160 is a single legal instruction. Selecting gcd(warp_col_tiles, 256) instead splits those into n32 / n16 atoms and gives up most of the tensor-core throughput.

Widening is restricted to multiples of 16 for two independent reasons: N % 16 == 8 already miscomputes for warp_cols > 1 (issue #2593), and the integer instruction tables in wgmma.h / wgmma_sp.h only instantiate multiples of 16 (plus 8 and 24), so a wider N % 16 == 8 would not even compile for the s8 path.

Parameters:

warp_col_tiles (int)

Return type:

int

class tilelang.cuda.intrinsics.macro.wgmma_macro_generator.WGMMADescriptorParams¶

Pre-computed WGMMA descriptor parameters, produced by compute_gmma_descriptor() (a port of CuTe make_gmma_desc) and consumed by init_wgmma_*_desc() and wgmma_*_atom().

swizzle_mode: tilelang.layout.SwizzleMode¶

Canonical swizzle mode; project to the descriptor field via wgmma_layout_type().

leading_byte_offset: int¶

LBO >> 4, ready to pass to T.initialize_wgmma_descriptor.

stride_byte_offset: int¶

SBO >> 4, ready to pass to T.initialize_wgmma_descriptor.

swizzle_atom_elems: int¶

Number of elements per swizzle atom along the non-K dimension.

k_atom_size: int¶

max(swizzle_atom_elems // micro_size_k, 1).

elems_in_bytes: int¶

DataType(dtype).bits // 8.

Type:

Byte width of a single element

is_k_major: bool¶

Whether the matrix is stored in K-major order (affects offset formula branching).

slice_byte_offset: object = 0¶

Physical byte offset (raw bytes) of the operand slice origin within its buffer; passed to T.increase_descriptor_offset after building the descriptor from the buffer base. 0 for a whole-buffer / base-origin operand. Computed by compute_gmma_descriptor() from the CuTe layout.

k_panel_elems: int | None = None¶

Element stride between consecutive K swizzle-atom panels (K-major only), read off the decoded layout by decode_k_panel_elems(). None only when no panel step is needed – see k_panel_stride(), which is how callers should read it.

k_panel_stride(mn_extent)¶

K-panel step for the atom offset formulas.

mn_extent is the operand’s own MN extent, used only for the single-panel / MN-major case where the ki // k_atom_size factor is always zero and any value works. Never reconstruct the multi-panel step from it: for a slice of a wider buffer the panels stay spaced by the buffer’s MN extent, not the operand’s.

Parameters:

mn_extent (int)

Return type:

int

tilelang.cuda.intrinsics.macro.wgmma_macro_generator.compute_gmma_descriptor(tl_layout, buffer, transposed, micro_size_k=16, region=None)¶

Port of CuTe make_gmma_desc.

Decode an arbitrary shared-memory tl_layout for buffer and compute the WGMMA descriptor parameters, accepting any WGMMA-canonical layout – not just the four “maker” layouts. A non-GMMA-canonical layout is a programming error, so the canonicity checks assert (mirroring CuTe’s ``static_assert``s).

transposed selects which logical axis is MN vs K (shape only): default [MN, K], transposed [K, MN]. The operand is required to be row-major, i.e. K-major iff not transposed; the contiguity detected from the layout (the GMMA mode owning the stride-1 sub-mode) is asserted to agree.

region is the operand’s per-axis ranges (one tvm.ir.Range per logical buffer mode), used to restrict the decoded layout to a sliced operand (e.g. B[:, j*64:...]). It may be None for a full-buffer operand or the atom-level API.

Parameters:
  • transposed (bool)

  • micro_size_k (int)

Return type:

WGMMADescriptorParams

class tilelang.cuda.intrinsics.macro.wgmma_macro_generator.TensorCoreIntrinEmitter(a_dtype='float16', b_dtype='float16', accum_dtype='float16', a_transposed=False, b_transposed=False, block_row_warps=2, block_col_warps=2, warp_row_tiles=8, warp_col_tiles=8, chunk=16, reduce_k=1, num_elems_per_byte=1, is_m_first=False, thread_var=None)¶

Bases: tilelang.cuda.intrinsics.macro.mma_macro_generator.TensorCoreIntrinEmitter

To eliminate Python syntax within TIR Macro.

Parameters:
  • a_dtype (str)

  • b_dtype (str)

  • accum_dtype (str)

  • a_transposed (bool)

  • b_transposed (bool)

  • block_row_warps (int)

  • block_col_warps (int)

  • warp_row_tiles (int)

  • warp_col_tiles (int)

  • chunk (int)

  • reduce_k (int)

  • num_elems_per_byte (int)

  • is_m_first (bool | None)

  • thread_var (tvm.tirx.Var | None)

wgmma_prefix: str¶
wgmma_inst_m: int¶
wgmma_inst_n: int¶
a_shared_layout: tilelang.layout.Layout = None¶
b_shared_layout: tilelang.layout.Layout = None¶
wgmma(A_region, B_region, C_region, clear_accum=False, wg_wait=0)¶
Parameters:
  • A_region (tvm.tirx.BufferRegion)

  • B_region (tvm.tirx.BufferRegion)

  • C_region (tvm.tirx.BufferRegion)

  • clear_accum (tvm.tirx.PrimExpr)

  • wg_wait (int)

wgmma_rs(A_region, B_region, C_region, clear_accum=False, wg_wait=0)¶
Parameters:
  • A_region (tvm.tirx.BufferRegion)

  • B_region (tvm.tirx.BufferRegion)

  • C_region (tvm.tirx.BufferRegion)

  • clear_accum (tvm.tirx.PrimExpr)

  • wg_wait (int)

property wgmma_num_inst_m: int¶

Number of WGMMA instruction atoms along the M dimension.

Return type:

int

property wgmma_num_inst_n: int¶

Number of WGMMA instruction atoms along the N dimension.

Return type:

int

property wgmma_num_k_atoms: int¶

Number of K-dimension micro-steps (chunk // micro_size_k).

Return type:

int

property wgmma_a_regs: int¶

Number of 32-bit registers occupied by the A fragment (RS variant).

Return type:

int

property wgmma_accum_regs: int¶

Number of 32-bit registers occupied by the accumulator fragment.

Return type:

int

compute_wgmma_b_desc_params(B_region)¶

Compute B descriptor parameters from the B shared buffer region.

Pure-Python helper (no TIR emitted); the returned WGMMADescriptorParams is consumed by init_wgmma_b_desc() and wgmma_*_atom().

Parameters:

B_region (tvm.tirx.BufferRegion)

Return type:

WGMMADescriptorParams

compute_wgmma_a_desc_params(A_region)¶

Compute A descriptor parameters from the A shared buffer region (SS variant).

Pure-Python helper (no TIR emitted); the returned WGMMADescriptorParams is consumed by init_wgmma_a_desc() and wgmma_ss_atom().

Parameters:

A_region (tvm.tirx.BufferRegion)

Return type:

WGMMADescriptorParams

init_wgmma_b_desc(desc_b, B_region, b_params)¶

Emit TIR to initialize a pre-allocated WGMMA B descriptor.

Parameters:
  • desc_b (Buffer) – A descriptor buffer allocated via T.alloc_wgmma_desc().

  • B_region (BufferRegion) – The B operand shared memory region.

  • b_params (WGMMADescriptorParams) – Pre-computed parameters from compute_wgmma_b_desc_params().

init_wgmma_a_desc(desc_a, A_region, a_params)¶

Emit TIR to initialize a pre-allocated WGMMA A descriptor (SS variant).

Parameters:
  • desc_a (Buffer) – A descriptor buffer allocated via T.alloc_wgmma_desc().

  • A_region (BufferRegion) – The A operand shared memory region.

  • a_params (WGMMADescriptorParams) – Pre-computed parameters from compute_wgmma_a_desc_params().

wgmma_fence_a(A_buf)¶

Emit warpgroup_fence_operand for the A fragment buffer.

Parameters:

A_buf (tvm.tirx.Buffer)

wgmma_fence_c(C_buf)¶

Emit warpgroup_fence_operand for the accumulator buffer.

Parameters:

C_buf (tvm.tirx.Buffer)

wgmma_arrive()¶

Emit warpgroup_arrive().

wgmma_commit()¶

Emit warpgroup_commit_batch().

wgmma_wait(n=0)¶

Emit warpgroup_wait(n).

Parameters:

n (int)

wgmma_rs_atom(A_buf, desc_b, C_buf, inst_m_idx, inst_n_idx, ki, b_params, clear_accum=False)¶

Emit a single WGMMA RS instruction for atom (inst_m_idx, inst_n_idx, ki).

Must be called between a wgmma_fence_a/wgmma_fence_c/wgmma_arrive sequence and a wgmma_commit/wgmma_wait sequence.

Calling this for every (j, i, ki) in T.grid(wgmma_num_inst_n, wgmma_num_inst_m, wgmma_num_k_atoms) produces identical TIR to wgmma_rs().

Parameters:
  • A_buf (Buffer) – Fragment buffer for operand A (in registers).

  • desc_b (Buffer) – Initialized B descriptor (from init_wgmma_b_desc).

  • C_buf (Buffer) – Accumulator fragment buffer.

  • inst_m_idx (int) – M-dimension atom index (0 .. wgmma_num_inst_m - 1).

  • inst_n_idx (int) – N-dimension atom index (0 .. wgmma_num_inst_n - 1).

  • ki (int) – K-dimension atom index (0 .. wgmma_num_k_atoms - 1).

  • b_params (WGMMADescriptorParams) – Pre-computed B descriptor parameters.

  • clear_accum (PrimExpr) – Whether to zero the accumulator on the first K atom.

wgmma_ss_atom(desc_a, desc_b, C_buf, inst_m_idx, inst_n_idx, ki, a_params, b_params, clear_accum=False)¶

Emit a single WGMMA SS instruction for atom (inst_m_idx, inst_n_idx, ki).

Must be called between fence/arrive and commit/wait sequences.

Parameters:
  • desc_a (Buffer) – Initialized A descriptor (from init_wgmma_a_desc).

  • desc_b (Buffer) – Initialized B descriptor (from init_wgmma_b_desc).

  • C_buf (Buffer) – Accumulator fragment buffer.

  • inst_m_idx (int) – M-dimension atom index (0 .. wgmma_num_inst_m - 1).

  • inst_n_idx (int) – N-dimension atom index (0 .. wgmma_num_inst_n - 1).

  • ki (int) – K-dimension atom index (0 .. wgmma_num_k_atoms - 1).

  • a_params (WGMMADescriptorParams) – Pre-computed A descriptor parameters.

  • b_params (WGMMADescriptorParams) – Pre-computed B descriptor parameters.

  • clear_accum (PrimExpr) – Whether to zero the accumulator on the first K atom.

make_mma_load_layout(local_buf, matrix='A')¶

Create a layout function for storing MMA results into a fragment buffer. This layout is used in conjunction with inverse_mma_store_layout to map fragment indices to threads and local indices.

Parameters:
  • local_buf (tir.Buffer) – The local buffer representing a fragment of a matrix.

  • matrix (str)

Returns:

A fragment object that describes how threads and indices in local_buf are laid out.

Return type:

T.Fragment

Raises:

AssertionError – If local_buf is not detected to be a fragment buffer.

make_mma_store_layout(local_buf)¶

Create a layout function for storing MMA results into a fragment buffer. This layout is used in conjunction with inverse_mma_store_layout to map fragment indices to threads and local indices.

Parameters:

local_buf (tir.Buffer) – The local buffer representing a fragment of a matrix.

Returns:

A fragment object that describes how threads and indices in local_buf are laid out.

Return type:

T.Fragment

Raises:

AssertionError – If local_buf is not detected to be a fragment buffer.