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:
BIN
src/kernels/addr_dump.co
Executable file
BIN
src/kernels/addr_dump.co
Executable file
Binary file not shown.
149
src/kernels/addr_dump.s
Normal file
149
src/kernels/addr_dump.s
Normal 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
BIN
src/kernels/coop_test.co
Executable file
Binary file not shown.
92
src/kernels/coop_test.s
Normal file
92
src/kernels/coop_test.s
Normal 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
BIN
src/kernels/lds_test.co
Executable file
Binary file not shown.
72
src/kernels/lds_test.s
Normal file
72
src/kernels/lds_test.s
Normal 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
BIN
src/kernels/matmul.co
Executable file
Binary file not shown.
279
src/kernels/matmul.s
Normal file
279
src/kernels/matmul.s
Normal 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
BIN
src/kernels/matmul_blocked.co
Executable file
Binary file not shown.
338
src/kernels/matmul_blocked.s
Normal file
338
src/kernels/matmul_blocked.s
Normal 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 barrier→compute 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
BIN
src/kernels/matmul_dbg.co
Executable file
Binary file not shown.
282
src/kernels/matmul_dbg.s
Normal file
282
src/kernels/matmul_dbg.s
Normal 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
BIN
src/kernels/matmul_small.co
Executable file
Binary file not shown.
310
src/kernels/matmul_small.s
Normal file
310
src/kernels/matmul_small.s
Normal 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
106
src/kernels/matvec.cl
Normal 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
BIN
src/kernels/matvec.co
Executable file
Binary file not shown.
138
src/kernels/matvec.s
Normal file
138
src/kernels/matvec.s
Normal 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
BIN
src/kernels/matvec_asm.co
Executable file
Binary file not shown.
BIN
src/kernels/superlinear.co
Executable file
BIN
src/kernels/superlinear.co
Executable file
Binary file not shown.
180
src/kernels/superlinear.s
Normal file
180
src/kernels/superlinear.s
Normal 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
BIN
src/kernels/test_store.co
Executable file
Binary file not shown.
41
src/kernels/test_store.s
Normal file
41
src/kernels/test_store.s
Normal 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
|
||||
Reference in New Issue
Block a user