tilelang.cuda.intrinsics.macro.wgmma_macro_generator¶
Attributes¶
Classes¶
Pre-computed WGMMA descriptor parameters, produced by |
|
To eliminate Python syntax within TIR Macro. |
Functions¶
|
Element stride between consecutive K swizzle-atom panels, off the layout. |
|
Widest legal WGMMA |
|
Port of CuTe |
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 asmn_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_modeis the operand tile’s K mode in element space (post-restrict). ReturnsNoneonly when no panel step is needed: an MN-major operand, or a layout with a single K panel, where theki // k_atom_sizefactor 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
Nthat tileswarp_col_tilesexactly.Hopper WGMMA accepts
Nin[8, 256]withN % 8 == 0, so an extent such as 96 or 160 is a single legal instruction. Selectinggcd(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 == 8already miscomputes forwarp_cols > 1(issue #2593), and the integer instruction tables inwgmma.h/wgmma_sp.honly instantiate multiples of 16 (plus 8 and 24), so a widerN % 16 == 8would 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 CuTemake_gmma_desc) and consumed byinit_wgmma_*_desc()andwgmma_*_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_offsetafter building the descriptor from the buffer base.0for a whole-buffer / base-origin operand. Computed bycompute_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().Noneonly when no panel step is needed – seek_panel_stride(), which is how callers should read it.
- k_panel_stride(mn_extent)¶
K-panel step for the atom offset formulas.
mn_extentis the operand’s own MN extent, used only for the single-panel / MN-major case where theki // k_atom_sizefactor 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_layoutforbufferand 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).transposedselects 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 iffnot transposed; the contiguity detected from the layout (the GMMA mode owning the stride-1 sub-mode) is asserted to agree.regionis the operand’s per-axis ranges (onetvm.ir.Rangeper logical buffer mode), used to restrict the decoded layout to a sliced operand (e.g.B[:, j*64:...]). It may beNonefor a full-buffer operand or the atom-level API.- Parameters:
transposed (bool)
micro_size_k (int)
- Return type:
- 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.TensorCoreIntrinEmitterTo eliminate Python syntax within TIR Macro.
- Parameters:
- wgmma_prefix: str¶
- wgmma_inst_m: int¶
- wgmma_inst_n: int¶
- 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
WGMMADescriptorParamsis consumed byinit_wgmma_b_desc()andwgmma_*_atom().- Parameters:
B_region (tvm.tirx.BufferRegion)
- Return type:
- 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
WGMMADescriptorParamsis consumed byinit_wgmma_a_desc()andwgmma_ss_atom().- Parameters:
A_region (tvm.tirx.BufferRegion)
- Return type:
- 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_operandfor the A fragment buffer.- Parameters:
A_buf (tvm.tirx.Buffer)
- wgmma_fence_c(C_buf)¶
Emit
warpgroup_fence_operandfor 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_arrivesequence and awgmma_commit/wgmma_waitsequence.Calling this for every
(j, i, ki)inT.grid(wgmma_num_inst_n, wgmma_num_inst_m, wgmma_num_k_atoms)produces identical TIR towgmma_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.