Skip to main content

gemm-optimization

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

설치로 이동

소스 정보

저장소
ROCm/FlyDSL
최근 소스 활동
2026년 8월 24일 06:42
감지된 SKILL.md 언어
영어
스타
280
포크
121

설치 방법

기본적으로 소스를 먼저 확인하는 Prompt가 선택됩니다. 직접 명령으로 전환하거나 로컬 사본을 다운로드할 수도 있습니다.

소스 파일 검토

설치 여부를 결정하기 전에 SKILL.md와 SkillsMP에 표시된 보조 파일을 읽어 보세요.

SKILL.md 표시 중

SKILL.md
소스 지침 · 읽기 전용 미리보기
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 ```python m_repeat = tile_m // 16 # M-direction 16x16 MFMA repeat count n_per_wave = tile_n // 4 # N range per wave (4 waves split tile_n) num_acc_n = n_per_wave // 16 # N-direction 16x16 accumulators per wave k_unroll = tile_k_bytes // a_elem_vec_pack // 64 # K-steps per tile (K64 micro-steps) ``` ### 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. ```python a_lds_elems = tile_m * tile_k # elements per A buffer @fx.struct class SharedStorage: a0: fx.Array[layout_elem, a_lds_elems, 16] # PONG buffer if lds_stage == 2: a1: fx.Array[layout_elem, a_lds_elems, 16] # PING buffer @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): ```python 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) # Phase 1: compute on PONG, prefetch to PING next_k1 = k_iv + tile_k store_a_tile_to_lds(prefetch_a_tile(next_k1), lds_a_ping) # A → PING LDS b_tile_ping = prefetch_b_tile(next_k1) # B → VGPR accs_in, _ = compute_tile(accs_in, b_tile_pong_in, lds_a_pong, a0_prefetch=a0pf_in) hot_loop_scheduler() # instruction hints rocdl.s_waitcnt(num_b_loads) gpu.barrier() a0_prefetch_ping = prefetch_a0_pack(lds_a_ping) # Phase 2: compute on PING, prefetch to PONG next_k2 = k_iv + (tile_k * 2) store_a_tile_to_lds(prefetch_a_tile(next_k2), lds_a_pong) # A → PONG LDS b_tile_pong_new = prefetch_b_tile(next_k2) # B → VGPR 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 ```python 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 ```python # In store_a_tile_to_lds(): 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 ```python # In lds_load_packs_k64(): 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 ```python a_regs = prefetch_a_tile(base_k) # buffer_load_dwordx4 → VGPR store_a_tile_to_lds(a_regs, lds_buffer) # ds_write from VGPR → LDS ``` **Asynchronous** (use_async_copy=True): Global → LDS directly ```python prefetch_a_to_lds(base_k, lds_buffer) # raw_ptr_buffer_load_lds (DMA) ``` 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: ```python b_tile = prefetch_b_tile(base_k) # buffer_load_dwordx4 → VGPR # b_tile structure: k_unroll × [(packs0[num_acc_n], packs1[num_acc_n])] ``` 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: ```python 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) ```python 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 # Prologue: pre-load first 2 LDS packs, interleave with first few MFMAs rocdl.sched_dsrd(2) # 2 ds_read for a0_prefetch rocdl.sched_mfma(1) rocdl.sched_mfma(1) # Main schedule: each iteration = 1 VMEM + mfma_group MFMAs + 1 ds_read + mfma_group MFMAs 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) # 1 global load (B tile or A tile) rocdl.sched_mfma(mfma_group) # N MFMA instructions rocdl.sched_dsrd(1) # 1 LDS read (A data) rocdl.sched_mfma(mfma_group) # N more MFMAs if sche_i >= dswr_start - 1: rocdl.sched_dswr(1) # LDS write (next A tile, tail end) rocdl.sched_barrier(0) # fence ``` ### 5.4 Key Scheduling Insights 1. **MFMA instructions dominate**: they form the backbone of the schedule 2. **LDS reads (ds_read) interleave with MFMAs**: one ds_read per 2×mfma_group MFMAs 3. **Global loads (VMEM) interleave**: one buffer_load per scheduler iteration 4. **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 5. **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: ```python dsrd_schedule = _build_scheduler(num_ds_load - dsrd_preload, mfma_total)
GitHub에서 보기
이 SKILL.md는 매우 커서 SkillsMP가 여기에는 첫 섹션만 미리 보여줍니다. GitHub에서 보기