File size: 2,704 Bytes
e92f76f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 | ; gfx942 MFMA GEMM Kernel Fragment: Direct LDS Staging Path
; Assumes: 16x16x16 MFMA, FP16 input, FP32 acc
; LDS allocation: A tile (0.5KB), B tile (0.5KB) ping-pong buffers
; s0-s3: A/B buffer descriptors (global mem)
; s4: K-loop counter
; s5: LDS base offset for A current tile
; s6: LDS base offset for B current tile
; s7: LDS base offset for A next tile (s5 + 0x200)
; s8: LDS base offset for B next tile (s6 + 0x200)
; v0-v3: Accumulator registers (c0-c3)
; v4-v7: A fragment registers
; v8-v11: B fragment registers
; ===== PROLOGUE: Load initial tiles into LDS[0] =====
buffer_load_lds v[0:1], s[0:3], 0 offen offset:0 lds:0 ; Load A tile (coalesced)
buffer_load_lds v[2:3], s[0:3], 0 offen offset:0 lds:0 ; Load B tile (coalesced)
s_waitcnt vmcnt(0) ; Wait for this wave's global loads
s_barrier ; Workgroup sync: all waves populated LDS[0]
; ===== MAIN K-LOOP =====
.L_loop:
; Prefetch NEXT tile into LDS[1] (overlap with current MFMA)
buffer_load_lds v[0:1], s[0:3], 0 offen offset:0 lds:1 ; A next
buffer_load_lds v[2:3], s[0:3], 0 offen offset:0 lds:1 ; B next
; Consume CURRENT tile from LDS[0] -> VGPR fragments
; (Example: 16x16 tile -> 4 lanes * 4 fragments each for MFMA)
ds_read_b32 v4, s5 offset:0 ; Lane 0: A frag0
ds_read_b32 v5, s5 offset:4 ; Lane 0: A frag1
ds_read_b32 v6, s5 offset:8 ; Lane 0: A frag2
ds_read_b32 v7, s5 offset:12 ; Lane 0: A frag3
ds_read_b32 v8, s6 offset:0 ; Lane 0: B frag0
ds_read_b32 v9, s6 offset:4 ; Lane 0: B frag1
ds_read_b32 v10, s6 offset:8 ; Lane 0: B frag2
ds_read_b32 v11, s6 offset:12 ; Lane 0: B frag3
; ... (other lanes implicitly handled by ds_read addressing)
s_waitcnt lgkmcnt(0) ; Wait for LDS reads to complete
; MFMA operation on VGPR-resident fragments
v_mfma_f32_16x16x16f16 v[0:3], v4, v5, v[0:3], 0, 0, 0 ; C += A*B
v_mfma_f32_16x16x16f16 v[0:3], v6, v7, v[0:3], 0, 0, 0
v_mfma_f32_16x16x16f16 v[0:3], v8, v9, v[0:3], 0, 0, 0
v_mfma_f32_16x16x16f16 v[0:3], v10, v11, v[0:3], 0, 0, 0
; Prepare for buffer swap: wait for next tile prefetch to finish
s_waitcnt vmcnt(0) ; Ensure global->LDS[1] done
s_barrier ; All waves agree: LDS[1] ready
; Swap ping-pong buffers (advance K pointers implicitly via s4)
s_add s5, s5, 0x400 ; A current = A next
s_add s6, s6, 0x400 ; B current = B next
s_sub s7, s7, 0x400 ; A next = A current (for next iter)
s_sub s8, s8, 0x400 ; B next = B current
s_sub s4, s4, 1 ; Decrement K tile counter
s_cbranch scc1 .L_loop ; Loop if more K tiles
; ===== EPILOGUE: Store C (not shown per focus on staging) =====
; v[0:3] holds final accumulators -> global store via vector_store |