Kernel development#
This tutorial covers advanced kernel development techniques in FlyDSL, including tiled data movement, MFMA instructions, shared memory, and performance optimization.
Tiled copies#
FlyDSL uses a hierarchical tiling model to partition data across blocks, warps, and threads:
import flydsl.compiler as flyc
import flydsl.expr as fx
@flyc.kernel
def copy_kernel(
A: fx.Tensor,
B: fx.Tensor,
):
tid = fx.thread_idx.x
bid = fx.block_idx.x
block_m = 8
block_n = 24
A = fx.rocdl.make_buffer_tensor(A)
B = fx.rocdl.make_buffer_tensor(B)
bA = fx.zipped_divide(A, (block_m, block_n))
bB = fx.zipped_divide(B, (block_m, block_n))
bA = fx.slice(bA, (None, bid))
bB = fx.slice(bB, (None, bid))
thr_layout = fx.make_layout((4, 1), (1, 1))
val_layout = fx.make_layout((1, 8), (1, 1))
copy_atom = fx.make_copy_atom(fx.rocdl.BufferCopy128b(), fx.Float32)
tile_mn, tv_layout = fx.make_layout_tv(thr_layout, val_layout)
tiled_copy = fx.make_tiled_copy(copy_atom, tv_layout, tile_mn)
thr_copy = tiled_copy.get_slice(tid)
partition_src = thr_copy.partition_S(bA)
partition_dst = thr_copy.partition_D(bB)
frag = fx.make_fragment_like(partition_src)
fx.copy(copy_atom, partition_src, frag)
fx.copy(copy_atom, frag, partition_dst)
See examples/02-tiledCopy.py for a complete working example.
MFMA instructions#
For matrix operations, FlyDSL supports AMD’s Matrix Fused Multiply-Add (MFMA)
instructions via make_mma_atom and make_tiled_mma:
import flydsl.compiler as flyc
import flydsl.expr as fx
block_m = 64
block_n = 64
block_k = 8
@flyc.kernel
def gemm_kernel(
A: fx.Tensor,
B: fx.Tensor,
C: fx.Tensor,
):
tid = fx.thread_idx.x
bid = fx.block_idx.x
A = fx.rocdl.make_buffer_tensor(A)
B = fx.rocdl.make_buffer_tensor(B)
C = fx.rocdl.make_buffer_tensor(C)
bA = fx.zipped_divide(A, (block_m, block_k))
bB = fx.zipped_divide(B, (block_n, block_k))
bC = fx.zipped_divide(C, (block_m, block_n))
bA = fx.slice(bA, (None, bid))
bB = fx.slice(bB, (None, bid))
bC = fx.slice(bC, (None, bid))
mma_atom = fx.make_mma_atom(fx.rocdl.MFMA(16, 16, 4, fx.Float32))
tiled_mma = fx.make_tiled_mma(mma_atom, fx.make_layout((2, 2, 1), (1, 2, 0)))
thr_mma = tiled_mma.thr_slice(tid)
copy_atom = fx.make_copy_atom(fx.rocdl.BufferCopy32b(), fx.Float32)
tiled_copy_A = fx.make_tiled_copy_A(copy_atom, tiled_mma)
tiled_copy_B = fx.make_tiled_copy_B(copy_atom, tiled_mma)
tiled_copy_C = fx.make_tiled_copy_C(copy_atom, tiled_mma)
thr_copy_A = tiled_copy_A.get_slice(tid)
thr_copy_B = tiled_copy_B.get_slice(tid)
thr_copy_C = tiled_copy_C.get_slice(tid)
copy_src_A = thr_copy_A.partition_S(bA)
copy_src_B = thr_copy_B.partition_S(bB)
copy_dst_C = thr_copy_C.partition_S(bC)
frag_A = thr_mma.make_fragment_A(bA)
frag_B = thr_mma.make_fragment_B(bB)
frag_C = thr_mma.make_fragment_C(bC)
copy_frag_A = thr_copy_A.retile(frag_A)
copy_frag_B = thr_copy_B.retile(frag_B)
copy_frag_C = thr_copy_C.retile(frag_C)
fx.copy(copy_atom, copy_src_A, copy_frag_A, pred=None)
fx.copy(copy_atom, copy_src_B, copy_frag_B, pred=None)
frag_C.fill(0)
fx.gemm(mma_atom, frag_C, frag_A, frag_B, frag_C)
fx.copy(copy_atom, copy_frag_C, copy_dst_C, pred=None)
See examples/03-tiledMma.py for a complete GEMM example and
kernels/gemm/preshuffle_gemm.py for a production GEMM implementation with
LDS pipeline.
Performance optimization#
Key optimization techniques demonstrated in the pre-built kernels:
LDS double-buffering: Overlap compute with data movement (
preshuffle_gemm)Buffer tensor operations: Hardware bounds-checked memory access (
fx.rocdl.make_buffer_tensor)Software pipelining: Hide memory latency with multi-stage pipelines
Pre-shuffled weights: Avoid runtime layout transformations for MFMA
Reference implementations#
Study these kernels for real-world patterns:
kernels/gemm/preshuffle_gemm.py– MFMA + LDS pipeline GEMMkernels/norm/softmax_kernel.py– online numerically stable softmaxkernels/norm/layernorm_kernel.py– fused normalizationkernels/attention/pa_decode_fp8.py– paged attention decode with FP8
See also
Kernel authoring guide – comprehensive kernel authoring reference
Pre-built kernel library guide – all pre-built kernels with configuration details
Testing & benchmarking guide – how to test and benchmark kernels