tilelang.metal.intrinsics.metal_macro_generator¶

Attributes¶

Classes¶

MPSIntrinEmitter

Metal simdgroup/cooperative tensor intrinsic emitter for GEMM operations.

Module Contents¶

tilelang.metal.intrinsics.metal_macro_generator.OPERAND_LEFT = 0¶
tilelang.metal.intrinsics.metal_macro_generator.OPERAND_RIGHT = 1¶
tilelang.metal.intrinsics.metal_macro_generator.OPERAND_DEST = 2¶
class tilelang.metal.intrinsics.metal_macro_generator.MPSIntrinEmitter(a_dtype='float16', b_dtype='float16', accum_dtype='float32', a_transposed=False, b_transposed=False, block_row_warps=1, block_col_warps=1, warp_row_tiles=8, warp_col_tiles=8, chunk=32, thread_var=None, a_stride_override=None, b_stride_override=None, inner_k_steps=1, use_cooperative_tensor=True)¶

Metal simdgroup/cooperative tensor intrinsic emitter for GEMM operations.

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)

  • thread_var (tvm.tirx.PrimExpr | None)

  • a_stride_override (int | None)

  • b_stride_override (int | None)

  • inner_k_steps (int)

  • use_cooperative_tensor (bool)

WARP_SIZE = 32¶
a_dtype = 'float16'¶
b_dtype = 'float16'¶
accum_dtype = 'float32'¶
a_transposed = False¶
b_transposed = False¶
block_row_warps = 1¶
block_col_warps = 1¶
warp_row_tiles = 8¶
warp_col_tiles = 8¶
chunk = 32¶
thread_var = None¶
a_stride_override = None¶
b_stride_override = None¶
inner_k_steps = 1¶
use_cooperative_tensor = True¶
warp_rows = 0¶
warp_cols = 0¶
get_thread_binding()¶

Return the thread index expression for the current kernel.

ldmatrix_a(A_local_buf, A_shared_buf, ki, k_inner=0)¶

Load matrix A tiles from memory into simdgroup/cooperative tensor buffers.

Parameters:
  • A_shared_buf (tvm.tirx.Buffer | tvm.tirx.BufferRegion)

  • k_inner (int)

ldmatrix_b(B_local_buf, B_shared_buf, ki, k_inner=0)¶

Load matrix B tiles from memory into simdgroup/cooperative tensor buffers.

Parameters:
  • B_shared_buf (tvm.tirx.Buffer | tvm.tirx.BufferRegion)

  • k_inner (int)

mma(A_local_buf, B_local_buf, C_local_buf, k_inner=0)¶

Perform matrix multiply-accumulate: C += A * B.

Parameters:

k_inner (int)

simdgroup_copy(C_simd_buf, C_dst, is_store=True)¶

Copy between register-backed Metal matrix buffers and memory.

make_cooperative_tensor_store_layout(local_buf)¶
simd_store(C_simd_buf, C_dst)¶

Store simdgroup/cooperative tensor local buffer to memory.

simd_load(C_simd_buf, C_src)¶

Load memory into simdgroup/cooperative tensor local buffer.