| name | gemm-optimization |
| description | Comprehensive guide to optimizing GEMM (General Matrix Multiply) kernels in FlyDSL on AMD CDNA GPUs. Covers tiling strategy, LDS ping-pong double-buffer, XOR bank-conflict swizzle, A/B data prefetch pipeline, 2-stage software pipelining, MFMA instruction scheduling (hot_loop_scheduler), epilogue strategies (direct store vs CShuffle), TFLOPS/bandwidth calculation, main-loop instruction count analysis, and bottleneck identification from ATT traces. Based on the production preshuffle_gemm kernel. Usage: /gemm-optimization
|
| allowed-tools | Read Edit Bash Grep Glob Agent |
GEMM Optimization Guide
Comprehensive guide to writing and optimizing high-performance GEMM kernels in
FlyDSL on AMD CDNA GPUs (MI300X gfx942, MI350 gfx950).
Based on the production kernels/gemm/preshuffle_gemm.py implementation.
1. Tiling Strategy
1.1 Three-Level Tiling
GEMM tiles the output C[M, N] and the reduction K into blocks:
C[M, N] = A[M, K] × B[K, N]^T
Grid mapping:
block_x → M tiles (tile_m rows each)
block_y → N tiles (tile_n cols each)
Thread mapping (256 threads = 4 waves × 64 lanes):
wave_id = tid // 64 ∈ [0, 3] → N dimension partitioning
lane_id = tid % 64 ∈ [0, 63] → M + N dimension within wave
lane_div_16 = lane_id // 16 → M dimension (4 groups of 16)
lane_mod_16 = lane_id % 16 → N dimension within MFMA
1.2 Derived Tile Parameters
m_repeat = tile_m // 16
n_per_wave = tile_n // 4
num_acc_n = n_per_wave // 16
k_unroll = tile_k_bytes // a_elem_vec_pack // 64
1.3 Recommended Tile Configurations
| Scenario | tile_m | tile_n | tile_k | Data Type | Notes |
|---|
| Small batch (M ≤ 32) | 16 | 64-128 | 256-512 | FP8/INT8 | Memory-bound, large tile_k for reuse |
| Medium batch | 64 | 256 | 128 | FP8/INT8/BF16 | Balanced compute/memory |
| Large batch (M ≥ 4096) | 128 | 256 | 128 | FP8/INT8 | Compute-dense, needs async copy |
| FP4 (gfx950) | 32-64 | 128-256 | 256 | FP4 | MFMA_SCALE instructions |
1.4 Tile Size Constraints
tile_m must be multiple of 16 (MFMA M dimension)
tile_n must be multiple of 64 (4 waves × 16 N per MFMA)
tile_k * elem_bytes must be multiple of 64 (K64-byte micro-step)
tile_m * tile_k * elem_bytes should fit comfortably in LDS (64KB on gfx942, 160KB on gfx950)
- B matrix is pre-shuffled to
(N/16, K/64, 4, 16, kpack_bytes) layout — tile_k must divide K evenly
1.5 MFMA Count Per Tile
Total MFMA instructions per tile:
MFMA_per_tile = k_unroll × m_repeat × num_acc_n × 2
↑ 2x K32 per K64 micro-step
Example (tile 64×256×128, FP8):
k_unroll = 128 / 64 = 2
m_repeat = 64 / 16 = 4
num_acc_n = 256 / 4 / 16 = 4
MFMA_per_tile = 2 × 4 × 4 × 2 = 64 MFMAs
Example (tile 64×256×512, FP8):
k_unroll = 512 / 64 = 8
MFMA_per_tile = 8 × 4 × 4 × 2 = 256 MFMAs
2. LDS Ping-Pong Double Buffer (2-Stage Pipeline)
2.1 Concept
With lds_stage=2, the kernel allocates two separate LDS buffers for the A
tile. While one buffer is used for MFMA computation, the next K-tile's A data
is loaded into the other buffer. This hides the global-to-LDS load latency.
Time →
Buffer PONG: [Compute tile_k=0] [ Load tile_k=2 ] [Compute tile_k=2] ...
Buffer PING: [ Load tile_k=1 ] [Compute tile_k=1] [ Load tile_k=3 ] ...
2.2 FlyDSL Implementation
Declare both A buffers as fx.Array fields of an @fx.struct and allocate them
with fx.SharedAllocator (the current LDS API — see kernels/gemm/preshuffle_gemm.py,
where a0/a1 are the pong/ping A buffers). In the default static=True mode the
compiler sizes the static LDS globals for you.
a_lds_elems = tile_m * tile_k
@fx.struct
class SharedStorage:
a0: fx.Array[layout_elem, a_lds_elems, 16]
if lds_stage == 2:
a1: fx.Array[layout_elem, a_lds_elems, 16]
@flyc.kernel
def kernel_gemm(...):
lds = fx.SharedAllocator().allocate(SharedStorage).peek()
lds_a_pong = lds.a0.view(fx.make_layout((tile_m, tile_k), (tile_k, 1)))
lds_a_ping = lds.a1.view(fx.make_layout((tile_m, tile_k), (tile_k, 1)))
The legacy flydsl.utils.smem_allocator.SmemAllocator path remains for un-migrated
kernels but is not recommended for new code.
2.3 Main Loop Structure (2-Stage)
Each iteration processes 2 K-tiles (one pong, one ping):
def _build_pingpong_body(k_iv, inner_state):
accs_in, bt_flat_in, a0pf_in = _unpack_state(inner_state)
b_tile_pong_in = _unflatten_b_tile(bt_flat_in)
next_k1 = k_iv + tile_k
store_a_tile_to_lds(prefetch_a_tile(next_k1), lds_a_ping)
b_tile_ping = prefetch_b_tile(next_k1)
accs_in, _ = compute_tile(accs_in, b_tile_pong_in, lds_a_pong,
a0_prefetch=a0pf_in)
hot_loop_scheduler()
rocdl.s_waitcnt(num_b_loads)
gpu.barrier()
a0_prefetch_ping = prefetch_a0_pack(lds_a_ping)
next_k2 = k_iv + (tile_k * 2)
store_a_tile_to_lds(prefetch_a_tile(next_k2), lds_a_pong)
b_tile_pong_new = prefetch_b_tile(next_k2)
accs_in, _ = compute_tile(accs_in, b_tile_ping, lds_a_ping,
a0_prefetch=a0_prefetch_ping)
hot_loop_scheduler()
rocdl.s_waitcnt(num_b_loads)
gpu.barrier()
a0_prefetch_pong_new = prefetch_a0_pack(lds_a_pong)
return _pack_state(accs_in, _flatten_b_tile(b_tile_pong_new),
a0_prefetch_pong_new)
2.4 LDS Size Budget
lds_tile_bytes = tile_m × tile_k × elem_bytes
2-stage total = 2 × lds_tile_bytes
+ CShuffle epilogue (optional): tile_m × tile_n × 2 bytes
Example (64×128, FP8): 2 × 64 × 128 = 16 KB total
Example (128×128, FP8): 2 × 128 × 128 = 32 KB total
Limits: 64 KB on gfx942, 160 KB on gfx950.
3. LDS XOR Bank-Conflict Swizzle
This section is the GEMM-specific implementation. For the general method --
diagnosing bank conflicts from ATT trace data, choosing swizzle vs padding, and
the gfx942 (32-bank) vs gfx950 (64-bank) differences -- see the
lds-optimization skill.
3.1 The Problem
A tile stored row-major in LDS with stride = tile_k creates bank conflicts when
multiple rows are read simultaneously (threads in the same wave access the same
bank for different addresses).
3.2 XOR Swizzle Formula
def swizzle_xor16(row, col, k_blocks16):
"""XOR-with-row swizzle at 16-byte granularity."""
rem = row % k_blocks16
return col ^ (rem * 16)
k_blocks16 = tile_k_bytes // a_elem_vec_pack // 16 — number of 16-byte blocks in K
- Applied to both write (global → LDS) and read (LDS → VGPR) paths
- Zero LDS overhead (no extra bytes), ~1 SALU instruction per address
3.3 Write Path
col_swz_bytes = swizzle_xor16(row_a_local, col_local_bytes, k_blocks16)
lds_offset = row_a_local * lds_stride_bytes + col_swz_bytes
lds_ptr.store(data, [lds_offset])
3.4 Read Path
col_base_swz_bytes = swizzle_xor16(curr_row_a_lds, col_base, k_blocks16)
lds_offset = curr_row_a_lds * lds_stride_bytes + col_base_swz_bytes
a_pack = lds_ptr.load([lds_offset])
Critical: swizzle must be consistent between write and read. If one path
uses swizzle but the other doesn't, data will be read from wrong positions.
4. Data Prefetch Pipeline
This section is the GEMM-specific pipeline. For the general transformation --
prologue, range(..., init=...) loop-carried state, epilogue -- see the
prefetch-data-load skill, and 10.2 below for the register budget.
4.1 A Matrix: Global → LDS
Two paths for loading A into LDS:
Synchronous (default): Global → VGPR → LDS
a_regs = prefetch_a_tile(base_k)
store_a_tile_to_lds(a_regs, lds_buffer)
Asynchronous (use_async_copy=True): Global → LDS directly
prefetch_a_to_lds(base_k, lds_buffer)
Async copy bypasses VGPR, reducing register pressure. Available on gfx942/gfx950.
4.2 B Matrix: Global → VGPR (Preshuffle)
B is pre-shuffled to match MFMA register layout, loaded directly to VGPR:
b_tile = prefetch_b_tile(base_k)
Each K64 micro-step needs 2 × num_acc_n i64 values for B (K32 × 2).
4.3 A0 Prefetch (Cross-Tile LDS Prefetch)
After gpu.barrier() completes (LDS is valid), immediately load the first A
pack from LDS into VGPR registers, overlapping with upcoming VMEM loads:
a0_prefetch = lds_load_packs_k64(row_a_lds, col_offset_base_bytes, lds_buffer)
This hides the first ds_read latency (~20-40 cycles) behind the global loads
that follow.
4.4 Pipeline Timeline
Iter i:
1. [VMEM] Load A(i+1) → PING LDS, Load B(i+1) → VGPR
2. [MFMA] Compute tile(i) using PONG LDS + B(i) VGPR
3. [SCHED] hot_loop_scheduler() — interleave MFMA with pending loads
4. [SYNC] s_waitcnt + barrier — wait for PING LDS to be valid
5. [LDS] A0 prefetch from PING — ds_read first pack
Swap PING ↔ PONG, repeat for i+1
5. Instruction Scheduling (hot_loop_scheduler)
5.1 Purpose
The hot_loop_scheduler() inserts rocdl.sched_* hints between the MFMA
compute phase and the next iteration's loads. These hints tell the compiler
how to interleave different instruction types to maximize pipeline utilization.
5.2 Scheduling Primitives
| Hint | Meaning | Maps to |
|---|
rocdl.sched_barrier(0) | Full scheduling barrier — no reordering across | Compiler fence |
rocdl.sched_mfma(N) | Allow N MFMA instructions | v_mfma_* |
rocdl.sched_dsrd(N) | Allow N LDS read instructions | ds_read_* |
rocdl.sched_dswr(N) | Allow N LDS write instructions | ds_write_* |
rocdl.sched_vmem(N) | Allow N global memory instructions | buffer_load_* |
5.3 Standard Schedule Pattern (gfx942, sync copy)
def hot_loop_scheduler():
mfma_group = num_acc_n
mfma_total = (k_unroll * 2) * m_repeat * mfma_group
mfma_per_iter = 2 * mfma_group
sche_iters = mfma_total // mfma_per_iter
rocdl.sched_dsrd(2)
rocdl.sched_mfma(1)
rocdl.sched_mfma(1)
dswr_tail = num_a_loads
dswr_start = max(sche_iters - dswr_tail - 2, 0)
for sche_i in range_constexpr(sche_iters):
rocdl.sched_vmem(1)
rocdl.sched_mfma(mfma_group)
rocdl.sched_dsrd(1)
rocdl.sched_mfma(mfma_group)
if sche_i >= dswr_start - 1:
rocdl.sched_dswr(1)
rocdl.sched_barrier(0)
5.4 Key Scheduling Insights
- MFMA instructions dominate: they form the backbone of the schedule
- LDS reads (ds_read) interleave with MFMAs: one ds_read per 2×mfma_group MFMAs
- Global loads (VMEM) interleave: one buffer_load per scheduler iteration
- LDS writes (ds_write) go at the tail: they overlap with the last MFMAs
of the current tile, landing before the
gpu.barrier() at iteration boundary
- dswr_start ensures LDS writes are scheduled early enough to complete
before the barrier, but late enough to not interfere with compute
5.5 Async Copy Schedule (gfx950)
For async copy, the scheduler uses _build_scheduler() to evenly distribute
ds_read and VMEM loads across all MFMAs:
dsrd_schedule = _build_scheduler(num_ds_load - dsrd_preload, mfma_total)