Initial extraction from isis: direct /dev/kfd GPU compute

AMD RDNA3 GPU driver via raw KFD ioctls. No ROCm, no OpenCL.

- Device discovery and topology queries
- VRAM/GTT/Userptr memory with shared address space
- PM4 compute queue and kernel dispatch
- Hand-written RDNA3 ASM kernels: matmul (3100 GFLOP/s), matvec, superlinear
- Tile-safe buffer padding for OOB protection
- Event-based interrupt wait (no CPU polling)

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
2026-04-06 20:38:10 +07:00
commit 66a9edb898
75 changed files with 5858 additions and 0 deletions

BIN
src/kernels/addr_dump.co Executable file

Binary file not shown.

149
src/kernels/addr_dump.s Normal file
View File

@@ -0,0 +1,149 @@
// Debug: dump computed addresses to Y buffer
// Y[lid*8 + 0] = v9 (W offset)
// Y[lid*8 + 1] = v11 (X offset)
// Y[lid*8 + 2] = s2 (W ptr lo)
// Y[lid*8 + 3] = s3 (W ptr hi)
// Y[lid*8 + 4] = s6 (X ptr lo)
// Y[lid*8 + 5] = s7 (X ptr hi)
// Y[lid*8 + 6] = s11 (K)
// Y[lid*8 + 7] = exec_lo
.amdgcn_target "amdgcn-amd-amdhsa--gfx1102"
.amdhsa_code_object_version 5
.text
.globl addr_dump
.p2align 8
.type addr_dump, @function
addr_dump:
s_mov_b32 s20, s2
s_load_b64 s[2:3], s[0:1], 0x00
s_load_b64 s[4:5], s[0:1], 0x08
s_load_b64 s[6:7], s[0:1], 0x10
s_load_b64 s[8:9], s[0:1], 0x18
s_load_b64 s[10:11], s[0:1], 0x20
s_load_b32 s12, s[0:1], 0x28
s_waitcnt lgkmcnt(0)
s_add_u32 s13, s10, 31
s_lshr_b32 s13, s13, 5
s_mov_b32 s14, 0
s_mov_b32 s15, s20
.Ldiv:
s_cmp_lt_u32 s15, s13
s_cbranch_scc1 .Ldiv_done
s_sub_u32 s15, s15, s13
s_add_u32 s14, s14, 1
s_branch .Ldiv
.Ldiv_done:
v_and_b32 v1, 31, v0
v_lshrrev_b32 v2, 5, v0
s_lshl_b32 s16, s15, 5
s_lshl_b32 s17, s14, 3
v_add_nc_u32 v3, s16, v1
v_add_nc_u32 v4, s17, v2
// bias load (same as matmul)
v_lshlrev_b32 v5, 2, v3
v_mov_b32 v6, 0
v_cmp_lt_u32 vcc_lo, v3, s10
s_and_saveexec_b32 s18, vcc_lo
global_load_b32 v6, v5, s[4:5]
s_mov_b32 exec_lo, s18
s_waitcnt vmcnt(0)
// precompute (same as matmul)
v_lshrrev_b32 v7, 3, v0
v_and_b32 v8, 7, v0
v_lshlrev_b32 v8, 2, v8
v_add_nc_u32 v9, s16, v7
v_mul_lo_u32 v9, v9, s11
v_add_nc_u32 v9, v9, v8
v_lshlrev_b32 v9, 2, v9
v_lshlrev_b32 v10, 4, v0
v_add_nc_u32 v11, s17, v2
v_mul_lo_u32 v11, v11, s11
v_add_nc_u32 v11, v11, v1
v_lshlrev_b32 v11, 2, v11
// dump 16 values per thread: Y[lid*16..lid*16+15]
v_lshlrev_b32 v20, 6, v0 // lid * 64 (16 dwords * 4 bytes)
global_store_b32 v20, v9, s[8:9] offset:0 // [0] W offset
global_store_b32 v20, v11, s[8:9] offset:4 // [1] X offset
v_mov_b32 v21, s14
global_store_b32 v20, v21, s[8:9] offset:8 // [2] s14 (wg_n)
v_mov_b32 v21, s15
global_store_b32 v20, v21, s[8:9] offset:12 // [3] s15 (wg_m)
v_mov_b32 v21, s16
global_store_b32 v20, v21, s[8:9] offset:16 // [4] s16 (wg_m*32)
v_mov_b32 v21, s17
global_store_b32 v20, v21, s[8:9] offset:20 // [5] s17 (wg_n*8)
v_mov_b32 v21, s13
global_store_b32 v20, v21, s[8:9] offset:24 // [6] s13 (num_wg_m)
v_mov_b32 v21, s20
global_store_b32 v20, v21, s[8:9] offset:28 // [7] s20 (wg_id)
global_store_b32 v20, v1, s[8:9] offset:32 // [8] v1 (thread_m)
global_store_b32 v20, v2, s[8:9] offset:36 // [9] v2 (thread_n)
v_mov_b32 v21, s11
global_store_b32 v20, v21, s[8:9] offset:40 // [10] s11 (K)
v_mov_b32 v21, s10
global_store_b32 v20, v21, s[8:9] offset:44 // [11] s10 (M)
v_mov_b32 v21, s12
global_store_b32 v20, v21, s[8:9] offset:48 // [12] s12 (N)
// intermediate: v11 before lshlrev = (s17+v2)*s11 + v1
// let's compute it fresh
v_add_nc_u32 v21, s17, v2
global_store_b32 v20, v21, s[8:9] offset:52 // [13] s17+v2
v_mul_lo_u32 v21, v21, s11
global_store_b32 v20, v21, s[8:9] offset:56 // [14] (s17+v2)*K
v_add_nc_u32 v21, v21, v1
global_store_b32 v20, v21, s[8:9] offset:60 // [15] (s17+v2)*K + v1
s_waitcnt vmcnt(0)
s_endpgm
.rodata
.p2align 6
.amdhsa_kernel addr_dump
.amdhsa_group_segment_fixed_size 0
.amdhsa_private_segment_fixed_size 0
.amdhsa_kernarg_size 48
.amdhsa_user_sgpr_kernarg_segment_ptr 1
.amdhsa_system_sgpr_workgroup_id_x 1
.amdhsa_next_free_vgpr 26
.amdhsa_next_free_sgpr 21
.amdhsa_float_denorm_mode_32 3
.amdhsa_float_denorm_mode_16_64 3
.amdhsa_wavefront_size32 1
.amdhsa_system_vgpr_workitem_id 0
.amdhsa_ieee_mode 1
.end_amdhsa_kernel
.amdgpu_metadata
---
amdhsa.version: [ 1, 2 ]
amdhsa.kernels:
- .name: addr_dump
.symbol: addr_dump.kd
.kernarg_segment_size: 48
.group_segment_fixed_size: 0
.private_segment_fixed_size: 0
.kernarg_segment_align: 8
.wavefront_size: 32
.sgpr_count: 21
.vgpr_count: 26
.max_flat_workgroup_size: 256
.args:
- { .size: 8, .offset: 0, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 8, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 16, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 24, .value_kind: global_buffer, .address_space: global }
- { .size: 4, .offset: 32, .value_kind: by_value }
- { .size: 4, .offset: 36, .value_kind: by_value }
- { .size: 4, .offset: 40, .value_kind: by_value }
...
.end_amdgpu_metadata

BIN
src/kernels/coop_test.co Executable file

Binary file not shown.

92
src/kernels/coop_test.s Normal file
View File

@@ -0,0 +1,92 @@
// test cooperative W tile load to LDS
// each thread: load 4 floats from W, store to LDS, barrier, read own row back
// kernargs: W(u64) Y(u64) M(u32) K(u32)
// dispatch: 1 WG of 256 threads, computes for first 32x8 tile
.amdgcn_target "amdgcn-amd-amdhsa--gfx1102"
.amdhsa_code_object_version 5
.text
.globl coop_test
.p2align 8
.type coop_test, @function
coop_test:
s_load_b64 s[2:3], s[0:1], 0x00 // W
s_load_b64 s[4:5], s[0:1], 0x08 // Y
s_load_b64 s[6:7], s[0:1], 0x10 // M, K (s6=M, s7=K)
s_waitcnt lgkmcnt(0)
// v0 = lid
v_and_b32 v1, 31, v0 // thread_m = lid & 31
v_lshrrev_b32 v2, 5, v0 // thread_n = lid >> 5
// W coop load: tile_row = lid/8, tile_col = (lid%8)*4
v_lshrrev_b32 v3, 3, v0 // tile_row
v_and_b32 v4, 7, v0 // lid & 7
v_lshlrev_b32 v4, 2, v4 // tile_col = (lid&7)*4
// W global offset = (tile_row * K + tile_col) * 4
v_mul_lo_u32 v5, v3, s7 // tile_row * K
v_add_nc_u32 v5, v5, v4 // + tile_col
v_lshlrev_b32 v5, 2, v5 // * 4 bytes
// global load 4 floats
global_load_b128 v[6:9], v5, s[2:3]
s_waitcnt vmcnt(0)
// LDS store offset = lid * 16
v_lshlrev_b32 v10, 4, v0
ds_store_b128 v10, v[6:9]
s_waitcnt lgkmcnt(0)
s_barrier
// Now read back: thread (thread_m, thread_n) reads W[thread_m][0] from LDS
// LDS offset = thread_m * 32 * 4 = thread_m * 128
v_lshlrev_b32 v11, 7, v1 // thread_m * 128
ds_load_b32 v12, v11 // W[thread_m][0]
s_waitcnt lgkmcnt(0)
// Store to Y[lid] = v12
v_lshlrev_b32 v13, 2, v0 // lid * 4
global_store_b32 v13, v12, s[4:5]
s_waitcnt vmcnt(0)
s_endpgm
.rodata
.p2align 6
.amdhsa_kernel coop_test
.amdhsa_group_segment_fixed_size 5120
.amdhsa_private_segment_fixed_size 0
.amdhsa_kernarg_size 24
.amdhsa_user_sgpr_kernarg_segment_ptr 1
.amdhsa_system_sgpr_workgroup_id_x 1
.amdhsa_next_free_vgpr 16
.amdhsa_next_free_sgpr 8
.amdhsa_float_denorm_mode_32 3
.amdhsa_float_denorm_mode_16_64 3
.amdhsa_wavefront_size32 1
.amdhsa_system_vgpr_workitem_id 0
.amdhsa_ieee_mode 1
.end_amdhsa_kernel
.amdgpu_metadata
---
amdhsa.version: [ 1, 2 ]
amdhsa.kernels:
- .name: coop_test
.symbol: coop_test.kd
.kernarg_segment_size: 24
.group_segment_fixed_size: 5120
.private_segment_fixed_size: 0
.kernarg_segment_align: 8
.wavefront_size: 32
.sgpr_count: 8
.vgpr_count: 16
.max_flat_workgroup_size: 256
.args:
- { .size: 8, .offset: 0, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 8, .value_kind: global_buffer, .address_space: global }
- { .size: 4, .offset: 16, .value_kind: by_value }
- { .size: 4, .offset: 20, .value_kind: by_value }
...
.end_amdgpu_metadata

BIN
src/kernels/lds_test.co Executable file

Binary file not shown.

72
src/kernels/lds_test.s Normal file
View File

@@ -0,0 +1,72 @@
// minimal LDS test: each thread stores lid to LDS, reads back neighbor's value
// output[lid] = lid + 1 (wrapped) to verify LDS sharing works
// kernarg: Y(u64)
// dispatch: global=256, local=256
.amdgcn_target "amdgcn-amd-amdhsa--gfx1102"
.amdhsa_code_object_version 5
.text
.globl lds_test
.p2align 8
.type lds_test, @function
lds_test:
// s[0:1] = kernarg, s2 = wg_id (unused)
s_load_b64 s[2:3], s[0:1], 0x00 // Y ptr
s_waitcnt lgkmcnt(0)
// v0 = lid
// store lid to LDS[lid*4]
v_lshlrev_b32 v1, 2, v0 // lid * 4
v_mov_b32 v2, v0 // value = lid
ds_store_b32 v1, v2
s_waitcnt lgkmcnt(0)
s_barrier
// read neighbor: LDS[(lid+1)%256 * 4]
v_add_nc_u32 v3, v0, 1
v_and_b32 v3, 255, v3 // (lid+1) % 256
v_lshlrev_b32 v3, 2, v3
ds_load_b32 v4, v3
s_waitcnt lgkmcnt(0)
// store to Y[lid]
global_store_b32 v1, v4, s[2:3]
s_waitcnt vmcnt(0)
s_endpgm
.rodata
.p2align 6
.amdhsa_kernel lds_test
.amdhsa_group_segment_fixed_size 1024
.amdhsa_private_segment_fixed_size 0
.amdhsa_kernarg_size 8
.amdhsa_user_sgpr_kernarg_segment_ptr 1
.amdhsa_system_sgpr_workgroup_id_x 1
.amdhsa_next_free_vgpr 8
.amdhsa_next_free_sgpr 4
.amdhsa_float_denorm_mode_32 3
.amdhsa_float_denorm_mode_16_64 3
.amdhsa_wavefront_size32 1
.amdhsa_system_vgpr_workitem_id 0
.amdhsa_ieee_mode 1
.end_amdhsa_kernel
.amdgpu_metadata
---
amdhsa.version: [ 1, 2 ]
amdhsa.kernels:
- .name: lds_test
.symbol: lds_test.kd
.kernarg_segment_size: 8
.group_segment_fixed_size: 1024
.private_segment_fixed_size: 0
.kernarg_segment_align: 8
.wavefront_size: 32
.sgpr_count: 4
.vgpr_count: 8
.max_flat_workgroup_size: 256
.args:
- { .size: 8, .offset: 0, .value_kind: global_buffer, .address_space: global }
...
.end_amdgpu_metadata

BIN
src/kernels/matmul.co Executable file

Binary file not shown.

279
src/kernels/matmul.s Normal file
View File

@@ -0,0 +1,279 @@
// rdna3 LDS-tiled matmul: Y = X * W^T + B
//
// Each workgroup computes a 32x8 tile of Y.
// 256 threads/WG: thread_m = lid%32, thread_n = lid/32
// Tiles K in chunks of TK=32. Each tile iteration:
// 1. Cooperative load W[32][32] and X[8][32] to LDS
// 2. barrier
// 3. Each thread reads its W row and X row from LDS, does 32 FMAs
// 4. barrier
//
// LDS layout: W at 0 (4096B), X at 4096 (1024B), total 5120B
//
// Cooperative load assignment (W: 1024 elts / 256 threads = 4 each):
// tile_row = lid/8, tile_col = (lid%8)*4
// global_load_b128 loads W[tile_row][tile_col..+3]
// ds_store_b128 at LDS offset lid*16
//
// Cooperative load assignment (X: 256 elts / 256 threads = 1 each):
// x_row = lid/32, x_col = lid%32
// global_load_b32 loads X[x_row][x_col]
// ds_store_b32 at LDS offset 4096 + lid*4
//
// SGPR: s[0:1]=kernarg, s2=TGID_X
// kernargs (48B): W(u64) B(u64) X(u64) Y(u64) M(u32) K(u32) N(u32)
// dispatch: grid.x = ceil(M/32)*ceil(N/8)*256, block.x = 256
.amdgcn_target "amdgcn-amd-amdhsa--gfx1102"
.amdhsa_code_object_version 5
.set TM, 32
.set TN, 8
.set TK, 32
.set LDS_X, 4096 // W uses 0..4095, X uses 4096..5119
.set LDS_SZ, 5120
.text
.globl matmul
.p2align 8
.type matmul, @function
matmul:
s_mov_b32 s20, s2 // save wg_id
s_load_b64 s[2:3], s[0:1], 0x00 // W
s_load_b64 s[4:5], s[0:1], 0x08 // B
s_load_b64 s[6:7], s[0:1], 0x10 // X
s_load_b64 s[8:9], s[0:1], 0x18 // Y
s_load_b64 s[10:11], s[0:1], 0x20 // M, K
s_load_b32 s12, s[0:1], 0x28 // N
s_waitcnt lgkmcnt(0)
// num_wg_m = ceil(M/32)
s_add_u32 s13, s10, TM - 1
s_lshr_b32 s13, s13, 5
// wg_m = wg_id % num_wg_m, wg_n = wg_id / num_wg_m
s_mov_b32 s14, 0
s_mov_b32 s15, s20
.Ldiv:
s_cmp_lt_u32 s15, s13
s_cbranch_scc1 .Ldiv_done
s_sub_u32 s15, s15, s13
s_add_u32 s14, s14, 1
s_branch .Ldiv
.Ldiv_done:
// s15 = wg_m, s14 = wg_n
// thread decomp
v_and_b32 v1, 31, v0 // thread_m = lid & 31
v_lshrrev_b32 v2, 5, v0 // thread_n = lid >> 5
// global output coords
s_lshl_b32 s16, s15, 5 // wg_m * 32
s_lshl_b32 s17, s14, 3 // wg_n * 8
v_add_nc_u32 v3, s16, v1 // global_m
v_add_nc_u32 v4, s17, v2 // global_n
// load bias into accumulator
v_lshlrev_b32 v5, 2, v3 // global_m * 4
v_mov_b32 v6, 0
v_cmp_lt_u32 vcc_lo, v3, s10
s_and_saveexec_b32 s18, vcc_lo
global_load_b32 v6, v5, s[4:5]
s_mov_b32 exec_lo, s18
s_waitcnt vmcnt(0)
// ======== Precompute cooperative load offsets ========
// W coop: tile_row = lid/8, tile_col = (lid%8)*4
v_lshrrev_b32 v7, 3, v0 // tile_row = lid >> 3
v_and_b32 v8, 7, v0 // lid & 7
v_lshlrev_b32 v8, 2, v8 // tile_col = (lid & 7) * 4
// W global byte offset for k_tile=0:
// ((wg_m*32 + tile_row) * K + tile_col) * 4
v_add_nc_u32 v9, s16, v7 // wg_m*32 + tile_row
v_mul_lo_u32 v9, v9, s11 // * K
v_add_nc_u32 v9, v9, v8 // + tile_col
v_lshlrev_b32 v9, 2, v9 // * 4 bytes
// v9 = W global load voffset (running, += TK*4 each iter)
// W LDS store offset = lid * 16 (4 floats * 4 bytes)
v_lshlrev_b32 v10, 4, v0
// X coop: x_row = lid/32 = v2, x_col = lid%32 = v1
// X global byte offset for k_tile=0:
// ((wg_n*8 + x_row) * K + x_col) * 4
v_add_nc_u32 v11, s17, v2 // wg_n*8 + lid/32
v_mul_lo_u32 v11, v11, s11 // * K
v_add_nc_u32 v11, v11, v1 // + lid%32
v_lshlrev_b32 v11, 2, v11 // * 4 bytes
// v11 = X global load voffset (running, += TK*4 each iter)
// X LDS store offset = LDS_X + lid * 4
v_lshlrev_b32 v12, 2, v0
v_add_nc_u32 v12, LDS_X, v12
// Compute-phase LDS read bases
v_lshlrev_b32 v13, 7, v1 // W: thread_m * 128
v_lshlrev_b32 v14, 7, v2
v_add_nc_u32 v14, LDS_X, v14 // X: LDS_X + thread_n * 128
// ======== Tile loop over K ========
s_mov_b32 s18, 0 // k_tile = 0
.Ltile_loop:
s_cmp_ge_u32 s18, s11 // k_tile >= K?
s_cbranch_scc1 .Ltile_done
// Phase 1: cooperative global LDS
global_load_b128 v[15:18], v9, s[2:3] // W: 4 consecutive floats
global_load_b32 v19, v11, s[6:7] // X: 1 float
s_waitcnt vmcnt(0)
ds_store_b128 v10, v[15:18] // W LDS
ds_store_b32 v12, v19 // X LDS
s_waitcnt lgkmcnt(0)
s_barrier
// Phase 2: compute 32 FMAs from LDS, unrolled 4x per block (8 blocks)
// block 0: tk=0..3
ds_load_b128 v[15:18], v13 offset:0
ds_load_b128 v[20:23], v14 offset:0
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 1: tk=4..7
ds_load_b128 v[15:18], v13 offset:16
ds_load_b128 v[20:23], v14 offset:16
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 2: tk=8..11
ds_load_b128 v[15:18], v13 offset:32
ds_load_b128 v[20:23], v14 offset:32
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 3: tk=12..15
ds_load_b128 v[15:18], v13 offset:48
ds_load_b128 v[20:23], v14 offset:48
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 4: tk=16..19
ds_load_b128 v[15:18], v13 offset:64
ds_load_b128 v[20:23], v14 offset:64
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 5: tk=20..23
ds_load_b128 v[15:18], v13 offset:80
ds_load_b128 v[20:23], v14 offset:80
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 6: tk=24..27
ds_load_b128 v[15:18], v13 offset:96
ds_load_b128 v[20:23], v14 offset:96
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 7: tk=28..31
ds_load_b128 v[15:18], v13 offset:112
ds_load_b128 v[20:23], v14 offset:112
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
s_barrier
// advance offsets to next K tile
v_add_nc_u32 v9, v9, TK * 4 // W += 128 bytes
v_add_nc_u32 v11, v11, TK * 4 // X += 128 bytes
s_add_u32 s18, s18, TK
s_branch .Ltile_loop
.Ltile_done:
// store Y[global_n][global_m] with bounds check
v_cmp_lt_u32 vcc_lo, v3, s10 // global_m < M
v_cmp_lt_u32 s19, v4, s12 // global_n < N
s_and_b32 s19, vcc_lo, s19
s_and_saveexec_b32 s20, s19
s_cbranch_execz .Ldone
v_mul_lo_u32 v15, v4, s10 // global_n * M
v_add_nc_u32 v15, v15, v3 // + global_m
v_lshlrev_b32 v15, 2, v15 // * 4 bytes
global_store_b32 v15, v6, s[8:9]
s_waitcnt vmcnt(0)
.Ldone:
s_endpgm
// ======== Kernel descriptor ========
.rodata
.p2align 6
.amdhsa_kernel matmul
.amdhsa_group_segment_fixed_size LDS_SZ
.amdhsa_private_segment_fixed_size 0
.amdhsa_kernarg_size 48
.amdhsa_user_sgpr_kernarg_segment_ptr 1
.amdhsa_system_sgpr_workgroup_id_x 1
.amdhsa_next_free_vgpr 24
.amdhsa_next_free_sgpr 21
.amdhsa_float_denorm_mode_32 3
.amdhsa_float_denorm_mode_16_64 3
.amdhsa_wavefront_size32 1
.amdhsa_system_vgpr_workitem_id 0
.amdhsa_ieee_mode 1
.end_amdhsa_kernel
.amdgpu_metadata
---
amdhsa.version: [ 1, 2 ]
amdhsa.kernels:
- .name: matmul
.symbol: matmul.kd
.kernarg_segment_size: 48
.group_segment_fixed_size: 5120
.private_segment_fixed_size: 0
.kernarg_segment_align: 8
.wavefront_size: 32
.sgpr_count: 21
.vgpr_count: 24
.max_flat_workgroup_size: 256
.args:
- { .size: 8, .offset: 0, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 8, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 16, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 24, .value_kind: global_buffer, .address_space: global }
- { .size: 4, .offset: 32, .value_kind: by_value }
- { .size: 4, .offset: 36, .value_kind: by_value }
- { .size: 4, .offset: 40, .value_kind: by_value }
...
.end_amdgpu_metadata

BIN
src/kernels/matmul_blocked.co Executable file

Binary file not shown.

View File

@@ -0,0 +1,338 @@
// rdna3 register-blocked LDS matmul: Y = X * W^T + B
//
// 4x4 output tile per thread: each thread computes 16 elements.
// TM=128, TN=32, TK=8, 256 threads/WG.
// threads_m = 128/4 = 32, threads_n = 32/4 = 8
// thread_m = lid % 32, thread_n = lid / 32
//
// LDS layout (TRANSPOSED for vectorized reads):
// W: W[TK][TM] at offset 0 8*128*4 = 4096 bytes
// X: X[TK][TN] at offset 4096 8*32*4 = 1024 bytes
// Total: 5120 bytes
//
// Cooperative load (per thread, 4 W elements + 1 X element):
// W: 128*8 = 1024 elts / 256 threads = 4 per thread
// tile_row = lid/2, tile_col = (lid%2)*4
// global: W[wg_m*128 + tile_row][k_tile + tile_col .. +3]
// LDS (transposed): W[tile_col+i][tile_row] for i=0..3
// = 4 scattered ds_store_b32 (stride = TM*4 = 512 bytes)
// X: 32*8 = 256 elts / 256 threads = 1 per thread
// x_row = lid/8, x_col = lid%8
// global: X[wg_n*32 + x_row][k_tile + x_col]
// LDS (transposed): X[x_col][x_row] at LDS_X + x_col*TN*4 + x_row*4
// = 1 ds_store_b32
//
// Compute phase (per thread, per tk step):
// W: ds_load_b128 reads W[tk][thread_m*4 .. thread_m*4+3] (contiguous after transpose)
// X: ds_load_b128 reads X[tk][thread_n*4 .. thread_n*4+3] (contiguous after transpose)
// 16 FMAs: outer product of 4 W values * 4 X values
// 2 LDS reads per 16 FMAs = 8x better than non-blocked kernel
//
// dispatch: grid.x = ceil(M/128)*ceil(N/32), block.x = 256
.amdgcn_target "amdgcn-amd-amdhsa--gfx1102"
.amdhsa_code_object_version 5
.set TM, 128
.set TN, 32
.set TK, 8
.set BM, 4 // block size per thread in M
.set BN, 4 // block size per thread in N
.set LDS_X, 4096 // W: TK*TM*4 = 8*128*4 = 4096
.set LDS_SZ, 5120 // + X: TK*TN*4 = 8*32*4 = 1024
.text
.globl matmul_blocked
.p2align 8
.type matmul_blocked, @function
matmul_blocked:
s_mov_b32 s20, s2 // save wg_id
s_load_b64 s[2:3], s[0:1], 0x00 // W
s_load_b64 s[4:5], s[0:1], 0x08 // B
s_load_b64 s[6:7], s[0:1], 0x10 // X
s_load_b64 s[8:9], s[0:1], 0x18 // Y
s_load_b64 s[10:11], s[0:1], 0x20 // M, K
s_load_b32 s12, s[0:1], 0x28 // N
s_waitcnt lgkmcnt(0)
// num_wg_m = ceil(M/128)
s_add_u32 s13, s10, TM - 1
s_lshr_b32 s13, s13, 7 // /128
// wg_m = wg_id % num_wg_m, wg_n = wg_id / num_wg_m
s_mov_b32 s14, 0
s_mov_b32 s15, s20
.Ldiv:
s_cmp_lt_u32 s15, s13
s_cbranch_scc1 .Ldiv_done
s_sub_u32 s15, s15, s13
s_add_u32 s14, s14, 1
s_branch .Ldiv
.Ldiv_done:
// s15 = wg_m, s14 = wg_n
// thread decomp: 32 threads in M, 8 in N
v_and_b32 v1, 31, v0 // thread_m = lid & 31
v_lshrrev_b32 v2, 5, v0 // thread_n = lid >> 5
// global output coords (base of 4x4 block)
s_lshl_b32 s16, s15, 7 // wg_m * 128
s_lshl_b32 s17, s14, 5 // wg_n * 32
v_lshlrev_b32 v3, 2, v1 // thread_m * 4 = base_m offset
v_add_nc_u32 v3, s16, v3 // global_m_base = wg_m*128 + thread_m*4
v_lshlrev_b32 v4, 2, v2 // thread_n * 4 = base_n offset
v_add_nc_u32 v4, s17, v4 // global_n_base = wg_n*32 + thread_n*4
// ======== Initialize 16 accumulators with bias ========
// acc[i][j] for i=0..3, j=0..3 in v16..v31
// acc[i][j] = B[global_m_base + i]
// Load 4 bias values
v_mov_b32 v16, 0
v_mov_b32 v17, 0
v_mov_b32 v18, 0
v_mov_b32 v19, 0
v_mov_b32 v20, 0
v_mov_b32 v21, 0
v_mov_b32 v22, 0
v_mov_b32 v23, 0
v_mov_b32 v24, 0
v_mov_b32 v25, 0
v_mov_b32 v26, 0
v_mov_b32 v27, 0
v_mov_b32 v28, 0
v_mov_b32 v29, 0
v_mov_b32 v30, 0
v_mov_b32 v31, 0
// load B[global_m_base+0..3]
v_lshlrev_b32 v5, 2, v3 // global_m_base * 4
v_cmp_lt_u32 vcc_lo, v3, s10 // bounds check
s_and_saveexec_b32 s18, vcc_lo
global_load_b128 v[16:19], v5, s[4:5] // B[m+0..3] acc[0..3][0]
s_mov_b32 exec_lo, s18
s_waitcnt vmcnt(0)
// Copy bias to all 4 N columns: acc[i][j] = B[m+i] for j=0..3
v_mov_b32 v20, v16 // acc[0][1] = B[m+0]
v_mov_b32 v24, v16 // acc[0][2]
v_mov_b32 v28, v16 // acc[0][3]
v_mov_b32 v21, v17 // acc[1][1] = B[m+1]
v_mov_b32 v25, v17 // acc[1][2]
v_mov_b32 v29, v17 // acc[1][3]
v_mov_b32 v22, v18 // acc[2][1]
v_mov_b32 v26, v18 // acc[2][2]
v_mov_b32 v30, v18 // acc[2][3]
v_mov_b32 v23, v19 // acc[3][1]
v_mov_b32 v27, v19 // acc[3][2]
v_mov_b32 v31, v19 // acc[3][3]
// ======== Precompute cooperative load offsets ========
// W coop: tile_row = lid/2 (0..127), tile_col = (lid%2)*4 (0 or 4)
v_lshrrev_b32 v5, 1, v0 // tile_row = lid >> 1
v_and_b32 v6, 1, v0 // lid & 1
v_lshlrev_b32 v6, 2, v6 // tile_col = (lid&1)*4
// W global byte offset for k_tile=0:
// ((wg_m*128 + tile_row) * K + tile_col) * 4
v_add_nc_u32 v7, s16, v5 // wg_m*128 + tile_row
v_mul_lo_u32 v7, v7, s11 // * K
v_add_nc_u32 v7, v7, v6 // + tile_col
v_lshlrev_b32 v7, 2, v7 // * 4 bytes
// v7 = W global load voffset (running, += TK*4 each iter)
// W LDS store offsets (transposed): W[tile_col+i][tile_row]
// base = tile_col * TM * 4 + tile_row * 4
v_mul_lo_u32 v8, v6, TM // tile_col * 128
v_lshlrev_b32 v8, 2, v8 // * 4 bytes
v_lshlrev_b32 v9, 2, v5 // tile_row * 4
v_add_nc_u32 v8, v8, v9 // base LDS offset for W store
// stride between consecutive tile_col values = TM*4 = 512
// X coop: x_row = lid/8 (0..31), x_col = lid%8 (0..7)
v_lshrrev_b32 v9, 3, v0 // x_row = lid >> 3
v_and_b32 v10, 7, v0 // x_col = lid & 7
// X global byte offset for k_tile=0:
// ((wg_n*32 + x_row) * K + x_col) * 4
v_add_nc_u32 v11, s17, v9 // wg_n*32 + x_row
v_mul_lo_u32 v11, v11, s11 // * K
v_add_nc_u32 v11, v11, v10 // + x_col
v_lshlrev_b32 v11, 2, v11 // * 4 bytes
// v11 = X global load voffset (running, += TK*4 each iter)
// X LDS store offset (transposed): X[x_col][x_row]
// = LDS_X + x_col * TN * 4 + x_row * 4
v_mul_lo_u32 v12, v10, TN // x_col * 32
v_lshlrev_b32 v12, 2, v12 // * 4
v_lshlrev_b32 v13, 2, v9 // x_row * 4
v_add_nc_u32 v12, v12, v13
v_add_nc_u32 v12, LDS_X, v12 // + LDS_X
// Compute-phase LDS read bases (transposed layout)
// W[tk][thread_m*4+0..3]: base = tk * TM * 4 + thread_m * 4 * 4
// = tk * 512 + thread_m * 16
// We use offset for tk, so base = thread_m * 16
v_lshlrev_b32 v13, 4, v1 // thread_m * 16
// X[tk][thread_n*4+0..3]: base = LDS_X + tk * TN * 4 + thread_n * 4 * 4
// = LDS_X + tk * 128 + thread_n * 16
v_lshlrev_b32 v14, 4, v2 // thread_n * 16
v_add_nc_u32 v14, LDS_X, v14 // + LDS_X
// ======== Tile loop over K (with prefetch) ========
// Prefetch hides global memory latency by issuing the next tile's
// load during the current tile's compute phase.
// v[40:43] = prefetch W, v44 = prefetch X
// v[32:35] = LDS W read, v[36:39] = LDS X read
s_mov_b32 s18, 0 // k_tile = 0
// Prologue: load first tile into prefetch regs
global_load_b128 v[40:43], v7, s[2:3] // W tile 0
global_load_b32 v44, v11, s[6:7] // X tile 0
.Ltile_loop:
// Wait for current tile's global data (first iter: prologue, then: prefetch)
s_waitcnt vmcnt(0)
// Store prefetched data to LDS (transposed)
ds_store_b32 v8, v40 // W[tile_col+0][tile_row]
ds_store_b32 v8, v41 offset:512 // W[tile_col+1][tile_row]
ds_store_b32 v8, v42 offset:1024 // W[tile_col+2][tile_row]
ds_store_b32 v8, v43 offset:1536 // W[tile_col+3][tile_row]
ds_store_b32 v12, v44 // X[x_col][x_row]
s_waitcnt lgkmcnt(0)
s_barrier
// Compute from LDS 8 k-steps × 16 FMAs = 128 FMAs
// Prefetch issued AFTER first LDS reads to not stall the barriercompute path
// tk=0
ds_load_b128 v[32:35], v13 offset:0
ds_load_b128 v[36:39], v14 offset:0
s_waitcnt lgkmcnt(0)
v_fmac_f32 v16, v32, v36
v_fmac_f32 v17, v33, v36
v_fmac_f32 v18, v34, v36
v_fmac_f32 v19, v35, v36
v_fmac_f32 v20, v32, v37
v_fmac_f32 v21, v33, v37
v_fmac_f32 v22, v34, v37
v_fmac_f32 v23, v35, v37
v_fmac_f32 v24, v32, v38
v_fmac_f32 v25, v33, v38
v_fmac_f32 v26, v34, v38
v_fmac_f32 v27, v35, v38
v_fmac_f32 v28, v32, v39
v_fmac_f32 v29, v33, v39
v_fmac_f32 v30, v34, v39
v_fmac_f32 v31, v35, v39
// Issue prefetch after first tk step overlaps with tk=1..7 compute
v_add_nc_u32 v7, v7, TK * 4 // W global += 32 bytes
v_add_nc_u32 v11, v11, TK * 4 // X global += 32 bytes
s_add_u32 s18, s18, TK
global_load_b128 v[40:43], v7, s[2:3] // prefetch next W
global_load_b32 v44, v11, s[6:7] // prefetch next X
// tk=1..7
.irp TK_OFF, 512, 1024, 1536, 2048, 2560, 3072, 3584
ds_load_b128 v[32:35], v13 offset:\TK_OFF
ds_load_b128 v[36:39], v14 offset:(\TK_OFF / 4)
s_waitcnt lgkmcnt(0)
v_fmac_f32 v16, v32, v36
v_fmac_f32 v17, v33, v36
v_fmac_f32 v18, v34, v36
v_fmac_f32 v19, v35, v36
v_fmac_f32 v20, v32, v37
v_fmac_f32 v21, v33, v37
v_fmac_f32 v22, v34, v37
v_fmac_f32 v23, v35, v37
v_fmac_f32 v24, v32, v38
v_fmac_f32 v25, v33, v38
v_fmac_f32 v26, v34, v38
v_fmac_f32 v27, v35, v38
v_fmac_f32 v28, v32, v39
v_fmac_f32 v29, v33, v39
v_fmac_f32 v30, v34, v39
v_fmac_f32 v31, v35, v39
.endr
s_barrier
// Loop if prefetched tile is valid
s_cmp_lt_u32 s18, s11 // s18 < K?
s_cbranch_scc1 .Ltile_loop
.Ltile_done:
// ======== Store 16 output elements ========
// Y[global_n_base+j][global_m_base+i] for i=0..3, j=0..3
// Y layout: row-major, Y[n][m], stride = M
// offset = (global_n_base + j) * M + (global_m_base + i)
// Compute base offset: global_n_base * M + global_m_base
v_mul_lo_u32 v5, v4, s10 // global_n_base * M
v_add_nc_u32 v5, v5, v3 // + global_m_base
v_lshlrev_b32 v5, 2, v5 // * 4 bytes
// Store row j=0: acc[0..3][0] = v16..v19
global_store_b128 v5, v[16:19], s[8:9]
// Row j=1: offset += M*4
s_lshl_b32 s19, s10, 2 // M * 4
v_add_nc_u32 v5, v5, s19
global_store_b128 v5, v[20:23], s[8:9]
// Row j=2
v_add_nc_u32 v5, v5, s19
global_store_b128 v5, v[24:27], s[8:9]
// Row j=3
v_add_nc_u32 v5, v5, s19
global_store_b128 v5, v[28:31], s[8:9]
s_waitcnt vmcnt(0)
s_endpgm
// ======== Kernel descriptor ========
.rodata
.p2align 6
.amdhsa_kernel matmul_blocked
.amdhsa_group_segment_fixed_size LDS_SZ
.amdhsa_private_segment_fixed_size 0
.amdhsa_kernarg_size 48
.amdhsa_user_sgpr_kernarg_segment_ptr 1
.amdhsa_system_sgpr_workgroup_id_x 1
.amdhsa_next_free_vgpr 48
.amdhsa_next_free_sgpr 21
.amdhsa_float_denorm_mode_32 3
.amdhsa_float_denorm_mode_16_64 3
.amdhsa_wavefront_size32 1
.amdhsa_system_vgpr_workitem_id 0
.amdhsa_ieee_mode 1
.end_amdhsa_kernel
.amdgpu_metadata
---
amdhsa.version: [ 1, 2 ]
amdhsa.kernels:
- .name: matmul_blocked
.symbol: matmul_blocked.kd
.kernarg_segment_size: 48
.group_segment_fixed_size: 5120
.private_segment_fixed_size: 0
.kernarg_segment_align: 8
.wavefront_size: 32
.sgpr_count: 21
.vgpr_count: 48
.max_flat_workgroup_size: 256
.args:
- { .size: 8, .offset: 0, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 8, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 16, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 24, .value_kind: global_buffer, .address_space: global }
- { .size: 4, .offset: 32, .value_kind: by_value }
- { .size: 4, .offset: 36, .value_kind: by_value }
- { .size: 4, .offset: 40, .value_kind: by_value }
...
.end_amdgpu_metadata

BIN
src/kernels/matmul_dbg.co Executable file

Binary file not shown.

282
src/kernels/matmul_dbg.s Normal file
View File

@@ -0,0 +1,282 @@
// rdna3 LDS-tiled matmul: Y = X * W^T + B
//
// Each workgroup computes a 32x8 tile of Y.
// 256 threads/WG: thread_m = lid%32, thread_n = lid/32
// Tiles K in chunks of TK=32. Each tile iteration:
// 1. Cooperative load W[32][32] and X[8][32] to LDS
// 2. barrier
// 3. Each thread reads its W row and X row from LDS, does 32 FMAs
// 4. barrier
//
// LDS layout: W at 0 (4096B), X at 4096 (1024B), total 5120B
//
// Cooperative load assignment (W: 1024 elts / 256 threads = 4 each):
// tile_row = lid/8, tile_col = (lid%8)*4
// global_load_b128 loads W[tile_row][tile_col..+3]
// ds_store_b128 at LDS offset lid*16
//
// Cooperative load assignment (X: 256 elts / 256 threads = 1 each):
// x_row = lid/32, x_col = lid%32
// global_load_b32 loads X[x_row][x_col]
// ds_store_b32 at LDS offset 4096 + lid*4
//
// SGPR: s[0:1]=kernarg, s2=TGID_X
// kernargs (48B): W(u64) B(u64) X(u64) Y(u64) M(u32) K(u32) N(u32)
// dispatch: grid.x = ceil(M/32)*ceil(N/8)*256, block.x = 256
.amdgcn_target "amdgcn-amd-amdhsa--gfx1102"
.amdhsa_code_object_version 5
.set TM, 32
.set TN, 8
.set TK, 32
.set LDS_X, 4096 // W uses 0..4095, X uses 4096..5119
.set LDS_SZ, 5120
.text
.globl matmul_dbg
.p2align 8
.type matmul, @function
matmul_dbg:
s_mov_b32 s20, s2 // save wg_id
s_load_b64 s[2:3], s[0:1], 0x00 // W
s_load_b64 s[4:5], s[0:1], 0x08 // B
s_load_b64 s[6:7], s[0:1], 0x10 // X
s_load_b64 s[8:9], s[0:1], 0x18 // Y
s_load_b64 s[10:11], s[0:1], 0x20 // M, K
s_load_b32 s12, s[0:1], 0x28 // N
s_waitcnt lgkmcnt(0)
// num_wg_m = ceil(M/32)
s_add_u32 s13, s10, TM - 1
s_lshr_b32 s13, s13, 5
// wg_m = wg_id % num_wg_m, wg_n = wg_id / num_wg_m
s_mov_b32 s14, 0
s_mov_b32 s15, s20
.Ldiv:
s_cmp_lt_u32 s15, s13
s_cbranch_scc1 .Ldiv_done
s_sub_u32 s15, s15, s13
s_add_u32 s14, s14, 1
s_branch .Ldiv
.Ldiv_done:
// s15 = wg_m, s14 = wg_n
// thread decomp
v_and_b32 v1, 31, v0 // thread_m = lid & 31
v_lshrrev_b32 v2, 5, v0 // thread_n = lid >> 5
// global output coords
s_lshl_b32 s16, s15, 5 // wg_m * 32
s_lshl_b32 s17, s14, 3 // wg_n * 8
v_add_nc_u32 v3, s16, v1 // global_m
v_add_nc_u32 v4, s17, v2 // global_n
// load bias into accumulator
v_lshlrev_b32 v5, 2, v3 // global_m * 4
v_mov_b32 v6, 0
v_cmp_lt_u32 vcc_lo, v3, s10
s_and_saveexec_b32 s18, vcc_lo
global_load_b32 v6, v5, s[4:5]
s_mov_b32 exec_lo, s18
// s_waitcnt vmcnt(0) -- no global loads
// ======== Precompute cooperative load offsets ========
// W coop: tile_row = lid/8, tile_col = (lid%8)*4
v_lshrrev_b32 v7, 3, v0 // tile_row = lid >> 3
v_and_b32 v8, 7, v0 // lid & 7
v_lshlrev_b32 v8, 2, v8 // tile_col = (lid & 7) * 4
// W global byte offset for k_tile=0:
// ((wg_m*32 + tile_row) * K + tile_col) * 4
v_add_nc_u32 v9, s16, v7 // wg_m*32 + tile_row
v_mul_lo_u32 v9, v9, s11 // * K
v_add_nc_u32 v9, v9, v8 // + tile_col
v_lshlrev_b32 v9, 2, v9 // * 4 bytes
// v9 = W global load voffset (running, += TK*4 each iter)
// W LDS store offset = lid * 16 (4 floats * 4 bytes)
v_lshlrev_b32 v10, 4, v0
// X coop: x_row = lid/32 = v2, x_col = lid%32 = v1
// X global byte offset for k_tile=0:
// ((wg_n*8 + x_row) * K + x_col) * 4
v_add_nc_u32 v11, s17, v2 // wg_n*8 + lid/32
v_mul_lo_u32 v11, v11, s11 // * K
v_add_nc_u32 v11, v11, v1 // + lid%32
v_lshlrev_b32 v11, 2, v11 // * 4 bytes
// v11 = X global load voffset (running, += TK*4 each iter)
// X LDS store offset = LDS_X + lid * 4
v_lshlrev_b32 v12, 2, v0
v_add_nc_u32 v12, LDS_X, v12
// Compute-phase LDS read bases
v_lshlrev_b32 v13, 7, v1 // W: thread_m * 128
v_lshlrev_b32 v14, 7, v2
v_add_nc_u32 v14, LDS_X, v14 // X: LDS_X + thread_n * 128
// ======== Tile loop over K ========
s_mov_b32 s18, 0 // k_tile = 0
.Ltile_loop:
s_cmp_ge_u32 s18, s11 // k_tile >= K?
s_cbranch_scc1 .Ltile_done
// Phase 1: cooperative global LDS
v_mov_b32 v15, 1.0
v_mov_b32 v16, 1.0
v_mov_b32 v17, 1.0
v_mov_b32 v18, 1.0
v_mov_b32 v19, 1.0
// s_waitcnt vmcnt(0) -- no global loads
ds_store_b128 v10, v[15:18] // W LDS
ds_store_b32 v12, v19 // X LDS
s_waitcnt lgkmcnt(0)
s_barrier
// Phase 2: compute 32 FMAs from LDS, unrolled 4x per block (8 blocks)
// block 0: tk=0..3
ds_load_b128 v[15:18], v13 offset:0
ds_load_b128 v[20:23], v14 offset:0
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 1: tk=4..7
ds_load_b128 v[15:18], v13 offset:16
ds_load_b128 v[20:23], v14 offset:16
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 2: tk=8..11
ds_load_b128 v[15:18], v13 offset:32
ds_load_b128 v[20:23], v14 offset:32
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 3: tk=12..15
ds_load_b128 v[15:18], v13 offset:48
ds_load_b128 v[20:23], v14 offset:48
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 4: tk=16..19
ds_load_b128 v[15:18], v13 offset:64
ds_load_b128 v[20:23], v14 offset:64
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 5: tk=20..23
ds_load_b128 v[15:18], v13 offset:80
ds_load_b128 v[20:23], v14 offset:80
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 6: tk=24..27
ds_load_b128 v[15:18], v13 offset:96
ds_load_b128 v[20:23], v14 offset:96
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
// block 7: tk=28..31
ds_load_b128 v[15:18], v13 offset:112
ds_load_b128 v[20:23], v14 offset:112
s_waitcnt lgkmcnt(0)
v_fmac_f32 v6, v15, v20
v_fmac_f32 v6, v16, v21
v_fmac_f32 v6, v17, v22
v_fmac_f32 v6, v18, v23
s_barrier
// advance offsets to next K tile
v_add_nc_u32 v9, v9, TK * 4 // W += 128 bytes
v_add_nc_u32 v11, v11, TK * 4 // X += 128 bytes
s_add_u32 s18, s18, TK
s_branch .Ltile_loop
.Ltile_done:
// store Y[global_n][global_m] with bounds check
v_cmp_lt_u32 vcc_lo, v3, s10 // global_m < M
v_cmp_lt_u32 s19, v4, s12 // global_n < N
s_and_b32 s19, vcc_lo, s19
s_and_saveexec_b32 s20, s19
s_cbranch_execz .Ldone
v_mul_lo_u32 v15, v4, s10 // global_n * M
v_add_nc_u32 v15, v15, v3 // + global_m
v_lshlrev_b32 v15, 2, v15 // * 4 bytes
global_store_b32 v15, v6, s[8:9]
// s_waitcnt vmcnt(0) -- no global loads
.Ldone:
s_endpgm
// ======== Kernel descriptor ========
.rodata
.p2align 6
.amdhsa_kernel matmul_dbg
.amdhsa_group_segment_fixed_size LDS_SZ
.amdhsa_private_segment_fixed_size 0
.amdhsa_kernarg_size 48
.amdhsa_user_sgpr_kernarg_segment_ptr 1
.amdhsa_system_sgpr_workgroup_id_x 1
.amdhsa_next_free_vgpr 24
.amdhsa_next_free_sgpr 21
.amdhsa_float_denorm_mode_32 3
.amdhsa_float_denorm_mode_16_64 3
.amdhsa_wavefront_size32 1
.amdhsa_system_vgpr_workitem_id 0
.amdhsa_ieee_mode 1
.end_amdhsa_kernel
.amdgpu_metadata
---
amdhsa.version: [ 1, 2 ]
amdhsa.kernels:
- .name: matmul_dbg
.symbol: matmul_dbg.kd
.kernarg_segment_size: 48
.group_segment_fixed_size: 5120
.private_segment_fixed_size: 0
.kernarg_segment_align: 8
.wavefront_size: 32
.sgpr_count: 21
.vgpr_count: 24
.max_flat_workgroup_size: 256
.args:
- { .size: 8, .offset: 0, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 8, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 16, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 24, .value_kind: global_buffer, .address_space: global }
- { .size: 4, .offset: 32, .value_kind: by_value }
- { .size: 4, .offset: 36, .value_kind: by_value }
- { .size: 4, .offset: 40, .value_kind: by_value }
...
.end_amdgpu_metadata

BIN
src/kernels/matmul_small.co Executable file

Binary file not shown.

310
src/kernels/matmul_small.s Normal file
View File

@@ -0,0 +1,310 @@
// TM=32 TT2x2 TK=8 wave32, PLR, single-buffer LDS
.amdgcn_target "amdgcn-amd-amdhsa--gfx1102"
.amdhsa_code_object_version 5
.set TM, 32
.set TN, 32
.set TK, 8
.set LDS_X, 1024
.set LDS_SZ, 2048
.text
.globl matmul_small
.p2align 8
.type matmul_small, @function
matmul_small:
s_mov_b32 s20, s2
s_load_b64 s[2:3], s[0:1], 0x00
s_load_b64 s[4:5], s[0:1], 0x08
s_load_b64 s[6:7], s[0:1], 0x10
s_load_b64 s[8:9], s[0:1], 0x18
s_load_b64 s[10:11], s[0:1], 0x20
s_load_b32 s12, s[0:1], 0x28
s_waitcnt lgkmcnt(0)
s_add_u32 s13, s10, TM - 1
s_lshr_b32 s13, s13, 5
s_mov_b32 s14, 0
s_mov_b32 s15, s20
.Ldiv:
s_cmp_lt_u32 s15, s13
s_cbranch_scc1 .Ldiv_done
s_sub_u32 s15, s15, s13
s_add_u32 s14, s14, 1
s_branch .Ldiv
.Ldiv_done:
v_and_b32 v1, 15, v0
v_lshrrev_b32 v2, 4, v0
s_lshl_b32 s16, s15, 5
s_lshl_b32 s17, s14, 5
v_lshlrev_b32 v3, 1, v1
v_add_nc_u32 v3, s16, v3
v_lshlrev_b32 v4, 1, v2
v_add_nc_u32 v4, s17, v4
v_mov_b32 v16, 0
v_mov_b32 v17, 0
v_mov_b32 v18, 0
v_mov_b32 v19, 0
v_lshlrev_b32 v5, 2, v3
v_cmp_lt_u32 vcc_lo, v3, s10
s_and_saveexec_b32 s18, vcc_lo
global_load_b64 v[16:17], v5, s[4:5]
s_mov_b32 exec_lo, s18
s_waitcnt vmcnt(0)
v_mov_b32 v18, v16
v_mov_b32 v19, v17
// W coop (col-major)
v_lshrrev_b32 v5, 5, v0
v_and_b32 v6, 31, v0
v_mul_lo_u32 v7, v5, s10
v_add_nc_u32 v7, v7, s16
v_add_nc_u32 v7, v7, v6
v_lshlrev_b32 v7, 2, v7
v_mul_lo_u32 v8, v5, TM
v_add_nc_u32 v8, v8, v6
v_lshlrev_b32 v8, 2, v8
s_lshl_b32 s22, s10, 5
// X coop (coalesced)
v_lshrrev_b32 v5, 3, v0
v_and_b32 v6, 7, v0
v_add_nc_u32 v9, s17, v5
v_mul_lo_u32 v9, v9, s11
v_add_nc_u32 v9, v9, v6
v_lshlrev_b32 v9, 2, v9
v_mul_lo_u32 v10, v6, TN
v_add_nc_u32 v10, v10, v5
v_lshlrev_b32 v10, 2, v10
v_add_nc_u32 v10, LDS_X, v10
// LDS read bases
v_lshlrev_b32 v11, 3, v1
v_lshlrev_b32 v12, 3, v2
v_add_nc_u32 v12, LDS_X, v12
// Prologue
s_mov_b32 s18, 0
global_load_b32 v20, v7, s[2:3]
global_load_b32 v21, v9, s[6:7]
s_waitcnt vmcnt(0)
ds_store_b32 v8, v20
ds_store_b32 v10, v21
s_waitcnt lgkmcnt(0)
s_barrier
v_add_nc_u32 v7, v7, s22
v_add_nc_u32 v9, v9, TK * 4
s_add_u32 s18, s18, TK
s_cmp_ge_u32 s18, s11
s_cbranch_scc1 .Llast_tile
.Ltile_loop:
// Prefetch tk=0
ds_load_2addr_b32 v[22:23], v11 offset0:0 offset1:1
ds_load_2addr_b32 v[24:25], v12 offset0:0 offset1:1
global_load_b32 v20, v7, s[2:3]
global_load_b32 v21, v9, s[6:7]
s_waitcnt lgkmcnt(0)
// tk=0
ds_load_2addr_b32 v[26:27], v11 offset0:(1*32) offset1:(1*32+1)
ds_load_2addr_b32 v[28:29], v12 offset0:(1*32) offset1:(1*32+1)
s_setprio 1
v_fmac_f32 v16, v22, v24
v_fmac_f32 v17, v23, v24
v_fmac_f32 v18, v22, v25
v_fmac_f32 v19, v23, v25
s_setprio 0
// tk=1
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[22:23], v11 offset0:(2*32) offset1:(2*32+1)
ds_load_2addr_b32 v[24:25], v12 offset0:(2*32) offset1:(2*32+1)
s_setprio 1
v_fmac_f32 v16, v26, v28
v_fmac_f32 v17, v27, v28
v_fmac_f32 v18, v26, v29
v_fmac_f32 v19, v27, v29
s_setprio 0
// tk=2
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[26:27], v11 offset0:(3*32) offset1:(3*32+1)
ds_load_2addr_b32 v[28:29], v12 offset0:(3*32) offset1:(3*32+1)
s_setprio 1
v_fmac_f32 v16, v22, v24
v_fmac_f32 v17, v23, v24
v_fmac_f32 v18, v22, v25
v_fmac_f32 v19, v23, v25
s_setprio 0
// tk=3
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[22:23], v11 offset0:(4*32) offset1:(4*32+1)
ds_load_2addr_b32 v[24:25], v12 offset0:(4*32) offset1:(4*32+1)
s_setprio 1
v_fmac_f32 v16, v26, v28
v_fmac_f32 v17, v27, v28
v_fmac_f32 v18, v26, v29
v_fmac_f32 v19, v27, v29
s_setprio 0
// tk=4
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[26:27], v11 offset0:(5*32) offset1:(5*32+1)
ds_load_2addr_b32 v[28:29], v12 offset0:(5*32) offset1:(5*32+1)
s_setprio 1
v_fmac_f32 v16, v22, v24
v_fmac_f32 v17, v23, v24
v_fmac_f32 v18, v22, v25
v_fmac_f32 v19, v23, v25
s_setprio 0
// tk=5
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[22:23], v11 offset0:(6*32) offset1:(6*32+1)
ds_load_2addr_b32 v[24:25], v12 offset0:(6*32) offset1:(6*32+1)
s_setprio 1
v_fmac_f32 v16, v26, v28
v_fmac_f32 v17, v27, v28
v_fmac_f32 v18, v26, v29
v_fmac_f32 v19, v27, v29
s_setprio 0
// tk=6
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[26:27], v11 offset0:(7*32) offset1:(7*32+1)
ds_load_2addr_b32 v[28:29], v12 offset0:(7*32) offset1:(7*32+1)
s_setprio 1
v_fmac_f32 v16, v22, v24
v_fmac_f32 v17, v23, v24
v_fmac_f32 v18, v22, v25
v_fmac_f32 v19, v23, v25
s_setprio 0
// tk=7
s_waitcnt lgkmcnt(0)
s_setprio 1
v_fmac_f32 v16, v26, v28
v_fmac_f32 v17, v27, v28
v_fmac_f32 v18, v26, v29
v_fmac_f32 v19, v27, v29
s_setprio 0
// Store next tile
s_waitcnt vmcnt(0)
ds_store_b32 v8, v20
ds_store_b32 v10, v21
s_waitcnt lgkmcnt(0)
s_barrier
v_add_nc_u32 v7, v7, s22
v_add_nc_u32 v9, v9, TK * 4
s_add_u32 s18, s18, TK
s_cmp_lt_u32 s18, s11
s_cbranch_scc1 .Ltile_loop
.Llast_tile:
ds_load_2addr_b32 v[22:23], v11 offset0:0 offset1:1
ds_load_2addr_b32 v[24:25], v12 offset0:0 offset1:1
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[26:27], v11 offset0:(1*32) offset1:(1*32+1)
ds_load_2addr_b32 v[28:29], v12 offset0:(1*32) offset1:(1*32+1)
s_setprio 1
v_fmac_f32 v16, v22, v24
v_fmac_f32 v17, v23, v24
v_fmac_f32 v18, v22, v25
v_fmac_f32 v19, v23, v25
s_setprio 0
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[22:23], v11 offset0:(2*32) offset1:(2*32+1)
ds_load_2addr_b32 v[24:25], v12 offset0:(2*32) offset1:(2*32+1)
s_setprio 1
v_fmac_f32 v16, v26, v28
v_fmac_f32 v17, v27, v28
v_fmac_f32 v18, v26, v29
v_fmac_f32 v19, v27, v29
s_setprio 0
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[26:27], v11 offset0:(3*32) offset1:(3*32+1)
ds_load_2addr_b32 v[28:29], v12 offset0:(3*32) offset1:(3*32+1)
s_setprio 1
v_fmac_f32 v16, v22, v24
v_fmac_f32 v17, v23, v24
v_fmac_f32 v18, v22, v25
v_fmac_f32 v19, v23, v25
s_setprio 0
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[22:23], v11 offset0:(4*32) offset1:(4*32+1)
ds_load_2addr_b32 v[24:25], v12 offset0:(4*32) offset1:(4*32+1)
s_setprio 1
v_fmac_f32 v16, v26, v28
v_fmac_f32 v17, v27, v28
v_fmac_f32 v18, v26, v29
v_fmac_f32 v19, v27, v29
s_setprio 0
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[26:27], v11 offset0:(5*32) offset1:(5*32+1)
ds_load_2addr_b32 v[28:29], v12 offset0:(5*32) offset1:(5*32+1)
s_setprio 1
v_fmac_f32 v16, v22, v24
v_fmac_f32 v17, v23, v24
v_fmac_f32 v18, v22, v25
v_fmac_f32 v19, v23, v25
s_setprio 0
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[22:23], v11 offset0:(6*32) offset1:(6*32+1)
ds_load_2addr_b32 v[24:25], v12 offset0:(6*32) offset1:(6*32+1)
s_setprio 1
v_fmac_f32 v16, v26, v28
v_fmac_f32 v17, v27, v28
v_fmac_f32 v18, v26, v29
v_fmac_f32 v19, v27, v29
s_setprio 0
s_waitcnt lgkmcnt(0)
ds_load_2addr_b32 v[26:27], v11 offset0:(7*32) offset1:(7*32+1)
ds_load_2addr_b32 v[28:29], v12 offset0:(7*32) offset1:(7*32+1)
s_setprio 1
v_fmac_f32 v16, v22, v24
v_fmac_f32 v17, v23, v24
v_fmac_f32 v18, v22, v25
v_fmac_f32 v19, v23, v25
s_setprio 0
s_waitcnt lgkmcnt(0)
s_setprio 1
v_fmac_f32 v16, v26, v28
v_fmac_f32 v17, v27, v28
v_fmac_f32 v18, v26, v29
v_fmac_f32 v19, v27, v29
s_setprio 0
// Store
v_mul_lo_u32 v5, v4, s10
v_add_nc_u32 v5, v5, v3
v_lshlrev_b32 v5, 2, v5
global_store_b64 v5, v[16:17], s[8:9]
s_lshl_b32 s19, s10, 2
v_add_nc_u32 v5, v5, s19
global_store_b64 v5, v[18:19], s[8:9]
s_waitcnt vmcnt(0)
s_endpgm
.rodata
.p2align 6
.amdhsa_kernel matmul_small
.amdhsa_group_segment_fixed_size LDS_SZ
.amdhsa_private_segment_fixed_size 0
.amdhsa_kernarg_size 48
.amdhsa_user_sgpr_kernarg_segment_ptr 1
.amdhsa_system_sgpr_workgroup_id_x 1
.amdhsa_next_free_vgpr 30
.amdhsa_next_free_sgpr 23
.amdhsa_float_denorm_mode_32 3
.amdhsa_float_denorm_mode_16_64 3
.amdhsa_wavefront_size32 1
.amdhsa_system_vgpr_workitem_id 0
.amdhsa_ieee_mode 1
.end_amdhsa_kernel
.amdgpu_metadata
---
amdhsa.version: [ 1, 2 ]
amdhsa.kernels:
- .name: matmul_small
.symbol: matmul_small.kd
.kernarg_segment_size: 48
.group_segment_fixed_size: 2048
.private_segment_fixed_size: 0
.kernarg_segment_align: 8
.wavefront_size: 32
.sgpr_count: 23
.vgpr_count: 30
.max_flat_workgroup_size: 256
.args:
- { .size: 8, .offset: 0, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 8, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 16, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 24, .value_kind: global_buffer, .address_space: global }
- { .size: 4, .offset: 32, .value_kind: by_value }
- { .size: 4, .offset: 36, .value_kind: by_value }
- { .size: 4, .offset: 40, .value_kind: by_value }
...
.end_amdgpu_metadata

106
src/kernels/matvec.cl Normal file
View File

@@ -0,0 +1,106 @@
// Matrix-vector multiply: y = W*x + b
// W is [out_dim x in_dim] row-major
// Each workgroup computes one output element
__kernel void matvec(
__global const float* W,
__global const float* b,
__global const float* x,
__global float* y,
uint out_dim,
uint in_dim
) {
uint row = get_global_id(0);
if (row >= out_dim) return;
float sum = b[row];
__global const float* w_row = W + row * in_dim;
// Vectorized inner loop
uint i = 0;
for (; i + 4 <= in_dim; i += 4) {
sum += w_row[i] * x[i];
sum += w_row[i+1] * x[i+1];
sum += w_row[i+2] * x[i+2];
sum += w_row[i+3] * x[i+3];
}
for (; i < in_dim; i++) {
sum += w_row[i] * x[i];
}
y[row] = sum;
}
// Fused synapse: matvec GLU SiLU layer_norm
// W is [out_dim*2 x in_dim], produces out_dim outputs
// scratch is [out_dim*2] temp space
__kernel void synapse_fused(
__global const float* W,
__global const float* b,
__global const float* x,
__global float* output,
__global float* scratch,
uint out_dim,
uint in_dim
) {
uint row = get_global_id(0);
uint total_rows = out_dim * 2;
if (row >= total_rows) return;
// Step 1: matvec into scratch
float sum = b[row];
__global const float* w_row = W + row * in_dim;
for (uint i = 0; i < in_dim; i++) {
sum += w_row[i] * x[i];
}
scratch[row] = sum;
barrier(CLK_GLOBAL_MEM_FENCE);
// Only first out_dim threads continue for GLU + SiLU
if (row >= out_dim) return;
// Step 2: GLU output[i] = scratch[i] * sigmoid(scratch[i + out_dim])
float val = scratch[row];
float gate = 1.0f / (1.0f + exp(-scratch[row + out_dim]));
float glu_out = val * gate;
// Step 3: SiLU x * sigmoid(x)
float silu_out = glu_out / (1.0f + exp(-glu_out));
output[row] = silu_out;
}
// GLU activation
__kernel void glu(
__global const float* input,
__global float* output,
uint half_dim
) {
uint i = get_global_id(0);
if (i >= half_dim) return;
float gate = 1.0f / (1.0f + exp(-input[i + half_dim]));
output[i] = input[i] * gate;
}
// SiLU (swish) in-place
__kernel void silu(
__global float* x,
uint n
) {
uint i = get_global_id(0);
if (i >= n) return;
float v = x[i];
x[i] = v / (1.0f + exp(-v));
}
// Elementwise add: y[i] += alpha * x[i]
__kernel void axpy(
__global float* y,
__global const float* x,
float alpha,
uint n
) {
uint i = get_global_id(0);
if (i >= n) return;
y[i] += alpha * x[i];
}

BIN
src/kernels/matvec.co Executable file

Binary file not shown.

138
src/kernels/matvec.s Normal file
View File

@@ -0,0 +1,138 @@
// rdna3 matvec kernel: y = W*x + b
// kernarg layout (40 bytes, no hidden args):
// +0x00: W pointer (u64) [out_dim x in_dim] row-major f32
// +0x08: b pointer (u64) [out_dim] f32 bias
// +0x10: x pointer (u64) [in_dim] f32 input
// +0x18: y pointer (u64) [out_dim] f32 output
// +0x20: out_dim (u32)
// +0x24: in_dim (u32)
//
// each workitem computes one output element (row).
// dispatch: global_size = out_dim, local_size = 1 (or up to 256)
// uses flat_load/flat_store with full 64-bit vgpr addresses (proven pattern)
.amdgcn_target "amdgcn-amd-amdhsa--gfx1102"
.amdhsa_code_object_version 5
.text
.globl matvec
.p2align 8
.type matvec, @function
matvec:
// s[0:1] = kernarg pointer
// v0 = workitem id
// load all kernargs
s_load_b64 s[2:3], s[0:1], 0x00 // W ptr
s_load_b64 s[4:5], s[0:1], 0x08 // b ptr
s_load_b64 s[6:7], s[0:1], 0x10 // x ptr
s_load_b64 s[8:9], s[0:1], 0x18 // y ptr
s_load_b64 s[10:11], s[0:1], 0x20 // out_dim(lo) | in_dim(hi)
s_waitcnt lgkmcnt(0)
// s10 = out_dim, s11 = in_dim (loaded as b64 from offset 0x20)
// bounds check: if v0 >= out_dim, skip
v_cmp_lt_u32 vcc_lo, v0, s10
s_and_saveexec_b32 s12, vcc_lo
s_cbranch_execz .Ldone
// row = v0
// row_byte_off = v0 * in_dim * 4 (byte offset into W for this row)
v_mul_lo_u32 v1, v0, s11 // v1 = row * in_dim (element offset)
v_lshlrev_b32 v1, 2, v1 // v1 = row * in_dim * 4 (byte offset)
// --- load bias b[row] using flat addressing ---
// addr = b_ptr + row*4
v_lshlrev_b32 v10, 2, v0 // v10 = row * 4
v_add_co_u32 v12, vcc_lo, s4, v10 // lo = b_lo + row*4
v_add_co_ci_u32 v13, vcc_lo, s5, 0, vcc_lo // hi with carry
flat_load_b32 v3, v[12:13] // v3 = b[row]
s_waitcnt vmcnt(0) lgkmcnt(0)
// v3 = accumulator (initialized to bias)
// loop over in_dim: sum += W[row*in_dim + i] * x[i]
s_mov_b32 s13, 0 // i = 0
.Lloop:
s_cmp_ge_u32 s13, s11 // i >= in_dim?
s_cbranch_scc1 .Lloop_done
// w_byte_off = row_byte_off + i*4
s_lshl_b32 s14, s13, 2 // s14 = i * 4
// --- load W[row][i] via flat ---
// addr = W_ptr + row_byte_off + i*4
v_add_nc_u32 v4, v1, s14 // v4 = row_byte_off + i*4
v_add_co_u32 v14, vcc_lo, s2, v4 // lo
v_add_co_ci_u32 v15, vcc_lo, s3, 0, vcc_lo // hi
flat_load_b32 v5, v[14:15] // v5 = W[row][i]
// --- load x[i] via flat ---
// addr = x_ptr + i*4
v_mov_b32 v6, s14 // v6 = i*4
v_add_co_u32 v16, vcc_lo, s6, v6 // lo
v_add_co_ci_u32 v17, vcc_lo, s7, 0, vcc_lo // hi
flat_load_b32 v7, v[16:17] // v7 = x[i]
s_waitcnt vmcnt(0) lgkmcnt(0)
// sum += W[row][i] * x[i]
v_fmac_f32 v3, v5, v7
// i++
s_add_u32 s13, s13, 1
s_branch .Lloop
.Lloop_done:
// --- store y[row] via flat ---
// addr = y_ptr + row*4
v_add_co_u32 v18, vcc_lo, s8, v10 // lo = y_lo + row*4
v_add_co_ci_u32 v19, vcc_lo, s9, 0, vcc_lo // hi
flat_store_b32 v[18:19], v3
s_waitcnt vmcnt(0) lgkmcnt(0)
.Ldone:
s_waitcnt vmcnt(0) lgkmcnt(0)
s_endpgm
// kernel descriptor
.rodata
.p2align 6
.amdhsa_kernel matvec
.amdhsa_group_segment_fixed_size 0
.amdhsa_private_segment_fixed_size 0
.amdhsa_kernarg_size 40
.amdhsa_user_sgpr_kernarg_segment_ptr 1
.amdhsa_next_free_vgpr 20
.amdhsa_next_free_sgpr 15
.amdhsa_float_denorm_mode_32 3
.amdhsa_float_denorm_mode_16_64 3
.amdhsa_wavefront_size32 1
.amdhsa_system_vgpr_workitem_id 0
.amdhsa_ieee_mode 1
.end_amdhsa_kernel
// AMDGPU metadata for HIP runtime module loading
.amdgpu_metadata
---
amdhsa.version: [ 1, 2 ]
amdhsa.kernels:
- .name: matvec
.symbol: matvec.kd
.kernarg_segment_size: 40
.group_segment_fixed_size: 0
.private_segment_fixed_size: 0
.kernarg_segment_align: 8
.wavefront_size: 32
.sgpr_count: 15
.vgpr_count: 20
.max_flat_workgroup_size: 256
.args:
- { .size: 8, .offset: 0, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 8, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 16, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 24, .value_kind: global_buffer, .address_space: global }
- { .size: 4, .offset: 32, .value_kind: by_value }
- { .size: 4, .offset: 36, .value_kind: by_value }
...
.end_amdgpu_metadata

BIN
src/kernels/matvec_asm.co Executable file

Binary file not shown.

BIN
src/kernels/superlinear.co Executable file

Binary file not shown.

180
src/kernels/superlinear.s Normal file
View File

@@ -0,0 +1,180 @@
// SuperLinear forward: N independent matvecs with per-neuron weights.
// Y[n*O+o] = B[n*O+o] + sum_k W[n*O*K + o*K + k] * X[n*K + k]
//
// Args: W(ptr), B(ptr), X(ptr), Y(ptr), N(u32), O(u32), K(u32)
// Grid: ceil(N*O / 256) workgroups, 256 threads each.
// Each thread computes one output element.
// K must be multiple of 4 and <= 32. Typical: K=4,8,16.
//
// Memory layout:
// W: [N * O * K] row-major (same as SuperLinear.weights)
// B: [N * O]
// X: [N * K] (trace, flat arena)
// Y: [N * O] (output)
.amdgcn_target "amdgcn-amd-amdhsa--gfx1102"
.amdhsa_code_object_version 5
.text
.globl superlinear_fwd
.p2align 8
.type superlinear_fwd, @function
superlinear_fwd:
// s[0:1] = kernarg_segment_ptr
// s2 = workgroup_id_x
// Load kernel args (7 args = 48 bytes)
s_load_b64 s[4:5], s[0:1], 0x00 // W ptr
s_load_b64 s[6:7], s[0:1], 0x08 // B ptr
s_load_b64 s[8:9], s[0:1], 0x10 // X ptr
s_load_b64 s[10:11], s[0:1], 0x18 // Y ptr
s_load_b32 s12, s[0:1], 0x20 // N (n_neurons)
s_load_b32 s13, s[0:1], 0x24 // O (out_per)
s_load_b32 s14, s[0:1], 0x28 // K (in_per)
s_waitcnt lgkmcnt(0)
// Global ID = workgroup_id * 256 + local_id
s_lshl_b32 s2, s2, 8 // s2 = workgroup_id * 256
v_add_nc_u32 v1, s2, v0 // v1 = global_id
// Total outputs = N * O
s_mul_i32 s15, s12, s13 // s15 = N * O
v_cmp_lt_u32 vcc_lo, v1, s15 // bounds check
s_and_saveexec_b32 s16, vcc_lo
s_cbranch_execz .Lexit
// Compute neuron = global_id / O, out_idx = global_id % O
// Use FP32 reciprocal for approximate division, then correct.
v_cvt_f32_u32 v4, s13 // v4 = float(O)
v_rcp_iflag_f32 v4, v4 // v4 = 1.0/O (approx)
v_cvt_f32_u32 v5, v1 // v5 = float(gid)
v_mul_f32 v5, v5, v4 // v5 = gid / O (float approx)
v_cvt_u32_f32 v2, v5 // v2 = neuron (truncated)
// Fix overshoot: if neuron * O > gid, decrement
v_mul_lo_u32 v4, v2, s13
v_cmp_gt_u32 vcc_lo, v4, v1
v_cndmask_b32 v5, 0, 1, vcc_lo
v_sub_nc_u32 v2, v2, v5
// Fix undershoot: if (neuron+1)*O <= gid, increment
v_add_nc_u32 v5, v2, 1
v_mul_lo_u32 v5, v5, s13
v_cmp_le_u32 vcc_lo, v5, v1
v_cndmask_b32 v5, 0, 1, vcc_lo
v_add_nc_u32 v2, v2, v5
v_mul_lo_u32 v4, v2, s13 // v4 = neuron * O
// out_idx = gid - neuron * O
v_sub_nc_u32 v3, v1, v4 // v3 = out_idx
// W byte offset = (neuron*O*K + out_idx*K) * 4
// = ((neuron*O + out_idx) * K) * 4
v_add_nc_u32 v5, v4, v3 // neuron*O + out_idx
v_mul_lo_u32 v5, v5, s14 // * K
v_lshlrev_b32 v5, 2, v5 // * 4 bytes
// X byte offset = neuron * K * 4
v_mul_lo_u32 v6, v2, s14
v_lshlrev_b32 v6, 2, v6
// B byte offset = (neuron*O + out_idx) * 4
v_add_nc_u32 v7, v4, v3
v_lshlrev_b32 v7, 2, v7
// Load bias
global_load_b32 v20, v7, s[6:7]
// Dot product: accumulate in v21
v_mov_b32 v21, 0 // acc = 0.0
// Loop over K in steps of 4 (vectorized loads)
s_mov_b32 s17, 0
.Ldot4_loop:
s_add_u32 s18, s17, 4
s_cmp_gt_u32 s18, s14 // if counter+4 > K, done with vec loop
s_cbranch_scc1 .Ldot4_done
// Load 4 floats from W and X
global_load_b128 v[12:15], v5, s[4:5] // W[0..3]
global_load_b128 v[16:19], v6, s[8:9] // X[0..3]
s_waitcnt vmcnt(0)
v_fmac_f32 v21, v12, v16
v_fmac_f32 v21, v13, v17
v_fmac_f32 v21, v14, v18
v_fmac_f32 v21, v15, v19
v_add_nc_u32 v5, v5, 16 // advance W ptr by 4 floats
v_add_nc_u32 v6, v6, 16 // advance X ptr by 4 floats
s_add_u32 s17, s17, 4
s_branch .Ldot4_loop
.Ldot4_done:
// Handle remaining 1-3 elements (scalar)
.Ldot1_loop:
s_cmp_ge_u32 s17, s14
s_cbranch_scc1 .Ldot1_done
global_load_b32 v12, v5, s[4:5]
global_load_b32 v13, v6, s[8:9]
s_waitcnt vmcnt(0)
v_fmac_f32 v21, v12, v13
v_add_nc_u32 v5, v5, 4
v_add_nc_u32 v6, v6, 4
s_add_u32 s17, s17, 1
s_branch .Ldot1_loop
.Ldot1_done:
// Y[gid] = acc + bias
s_waitcnt lgkmcnt(0)
v_add_f32 v21, v21, v20
// Store
v_lshlrev_b32 v1, 2, v1
global_store_b32 v1, v21, s[10:11]
.Lexit:
s_waitcnt vmcnt(0)
s_endpgm
// Kernel descriptor
.rodata
.p2align 6
.amdhsa_kernel superlinear_fwd
.amdhsa_group_segment_fixed_size 0
.amdhsa_private_segment_fixed_size 0
.amdhsa_kernarg_size 48
.amdhsa_user_sgpr_kernarg_segment_ptr 1
.amdhsa_system_sgpr_workgroup_id_x 1
.amdhsa_next_free_vgpr 22
.amdhsa_next_free_sgpr 19
.amdhsa_float_denorm_mode_32 3
.amdhsa_float_denorm_mode_16_64 3
.amdhsa_wavefront_size32 1
.amdhsa_system_vgpr_workitem_id 0
.amdhsa_ieee_mode 1
.end_amdhsa_kernel
.amdgpu_metadata
---
amdhsa.version: [ 1, 2 ]
amdhsa.kernels:
- .name: superlinear_fwd
.symbol: superlinear_fwd.kd
.kernarg_segment_size: 48
.group_segment_fixed_size: 0
.private_segment_fixed_size: 0
.kernarg_segment_align: 8
.wavefront_size: 32
.sgpr_count: 19
.vgpr_count: 22
.max_flat_workgroup_size: 256
.args:
- { .size: 8, .offset: 0, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 8, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 16, .value_kind: global_buffer, .address_space: global }
- { .size: 8, .offset: 24, .value_kind: global_buffer, .address_space: global }
- { .size: 4, .offset: 32, .value_kind: by_value }
- { .size: 4, .offset: 36, .value_kind: by_value }
- { .size: 4, .offset: 40, .value_kind: by_value }
...
.end_amdgpu_metadata

BIN
src/kernels/test_store.co Executable file

Binary file not shown.

41
src/kernels/test_store.s Normal file
View File

@@ -0,0 +1,41 @@
// absolute minimal test: store 42.0 to y[workitem_id]
// kernargs: y pointer (u64) at offset 0
.amdgcn_target "amdgcn-amd-amdhsa--gfx1102"
.amdhsa_code_object_version 5
.text
.globl test_store
.p2align 8
.type test_store, @function
test_store:
// s[0:1] = kernarg pointer
s_load_b64 s[2:3], s[0:1], 0x00 // y ptr
s_waitcnt lgkmcnt(0)
// build full 64-bit address in v[2:3] = s[2:3] + v0*4
v_lshlrev_b32 v1, 2, v0 // v1 = workitem_id * 4
v_add_co_u32 v2, vcc_lo, s2, v1 // v2 = y_lo + offset
v_add_co_ci_u32 v3, vcc_lo, s3, 0, vcc_lo // v3 = y_hi + carry
// store constant 42.0
v_mov_b32 v4, 0x42280000 // v4 = 42.0f
flat_store_b32 v[2:3], v4
s_waitcnt vmcnt(0) lgkmcnt(0)
s_endpgm
.rodata
.p2align 6
.amdhsa_kernel test_store
.amdhsa_group_segment_fixed_size 0
.amdhsa_private_segment_fixed_size 0
.amdhsa_kernarg_size 8
.amdhsa_user_sgpr_kernarg_segment_ptr 1
.amdhsa_next_free_vgpr 5
.amdhsa_next_free_sgpr 4
.amdhsa_float_denorm_mode_32 3
.amdhsa_float_denorm_mode_16_64 3
.amdhsa_wavefront_size32 1
.amdhsa_system_vgpr_workitem_id 0
.amdhsa_ieee_mode 1
.end_amdhsa_kernel