custom
code
sovereign-compute
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