code wiki / (root) / nx_dot_simd_demo.nx

nx_dot_simd_demo.nx source

↩ module page · 165 lines · 5827 B

1// nx_dot_simd_demo.nx -- parallel SIMD dot product demo. 2// 3// dot(a, b) = sum(a[i] * b[i]) for i in 0..N 4// 5// Three implementations, all bit-exact: 6// 1. scalar: trivial for-loop 7// 2. SIMD: i32x8 vmul + vreduce_sum per chunk; i64 accumulator 8// 3. parallel+SIMD: split chunks across pool workers, each 9// worker does SIMD inner loop, atomic FAA into shared acc 10// 11// Composes L7 nx_parallel + L8 SIMD i32x8 into a real ML-class 12// dot-product kernel. 13 14// nx_safety_envelope: 15// intended_use: AUTO_APPLIED -- primitive-specific tuning queued 16// sil_target: SIL1 17// evidence: [bulk_applied_2026-05-16, see-file-comment-for-detail] 18// verdict: NOT_YET_EVALUATED 19 20import "nx_kernel_v2.nx" 21import "nx_log.nx" 22import "nx_atom.nx" 23import "nx_clock.nx" 24import "nx_thread_pool.nx" 25import "nx_parallel.nx" 26import "nx_hw.nx" 27 28const N: i64 = 8192 29 30// Scalar reference -- treats each 4-byte slot as i32. 31func scalar_dot(a_i32: *u8, b_i32: *u8, n: i64) -> i64 { 32 var s: i64 = 0 33 var i: i64 = 0 34 while i < n { 35 let a_addr: i64 = (a_i32 as i64) + i * 4 36 let b_addr: i64 = (b_i32 as i64) + i * 4 37 // Load i32 by reading 8 bytes + masking low 32. 38 let a_p: *i64 = a_addr as *i64 39 let b_p: *i64 = b_addr as *i64 40 let av: i64 = *a_p & 0xFFFFFFFF 41 let bv: i64 = *b_p & 0xFFFFFFFF 42 s = s + av * bv 43 i = i + 1 44 } 45 return s 46} 47 48// SIMD i32x8 dot. Each chunk: vload+vload+vmul+vreduce, accumulate 49// scalar. Safe vs overflow (i64 acc throughout). 50func simd_dot(a_i32: *u8, b_i32: *u8, n: i64) -> i64 { 51 var s: i64 = 0 52 var i: i64 = 0 53 while i < n { 54 let a_p: *i64 = ((a_i32 as i64) + i * 4) as *i64 55 let b_p: *i64 = ((b_i32 as i64) + i * 4) as *i64 56 let va: i64 = __simd_vload_i32_x8(a_p) 57 let vb: i64 = __simd_vload_i32_x8(b_p) 58 let vp: i64 = __simd_vmul_i32_x8(va, vb) 59 let cs: i64 = __simd_vreduce_sum_i32_x8(vp) 60 s = s + cs 61 i = i + 8 62 } 63 return s 64} 65 66// Worker static context. Non-main helpers read/write these now 67// that the gp-relax PREVENT fix landed. 68static DOT_A_PTR: i64 69static DOT_B_PTR: i64 70static DOT_ACC_PTR: i64 71static DOT_CHUNK_SIZE: i64 72 73func _set_dot_ctx(a: i64, b: i64, acc: i64, chunk: i64) -> i64 { 74 DOT_A_PTR = a 75 DOT_B_PTR = b 76 DOT_ACC_PTR = acc 77 DOT_CHUNK_SIZE = chunk 78 return 0 79} 80 81func _read_dot_a() -> i64 { return DOT_A_PTR } 82func _read_dot_b() -> i64 { return DOT_B_PTR } 83func _read_dot_acc() -> i64 { return DOT_ACC_PTR } 84func _read_dot_chunk() -> i64 { return DOT_CHUNK_SIZE } 85 86// Per-worker SIMD dot kernel. chunk_idx selects the slice. 87func simd_dot_worker(chunk_idx: i64) -> i64 { 88 let chunk_size: i64 = _read_dot_chunk() 89 let start: i64 = chunk_idx * chunk_size 90 let end: i64 = start + chunk_size 91 let a_base: i64 = _read_dot_a() 92 let b_base: i64 = _read_dot_b() 93 let acc_ptr: *i64 = _read_dot_acc() as *i64 94 var local: i64 = 0 95 var i: i64 = start 96 while i < end { 97 let a_p: *i64 = (a_base + i * 4) as *i64 98 let b_p: *i64 = (b_base + i * 4) as *i64 99 let va: i64 = __simd_vload_i32_x8(a_p) 100 let vb: i64 = __simd_vload_i32_x8(b_p) 101 let vp: i64 = __simd_vmul_i32_x8(va, vb) 102 let cs: i64 = __simd_vreduce_sum_i32_x8(vp) 103 local = local + cs 104 i = i + 8 105 } 106 nx_atom_faa_i64(acc_ptr, local, NX_MO_SEQ_CST) 107 return 0 108} 109 110func main() -> nx_exit { 111 let a_raw: *u8 = sys_mmap(N * 4 + 64) 112 let b_raw: *u8 = sys_mmap(N * 4 + 64) 113 // a[i] = (i + 1) mod 1024; b[i] = ((i + 1) * 2) mod 1024 114 // (keep small to ensure i32 products fit and scalar matches SIMD) 115 var i: i64 = 0 116 while i < N { 117 let a_p: *i64 = ((a_raw as i64) + i * 4) as *i64 118 let b_p: *i64 = ((b_raw as i64) + i * 4) as *i64 119 let av: i64 = (i + 1) & 0x3FF 120 let bv: i64 = ((i + 1) * 2) & 0x3FF 121 *a_p = (*a_p & 0xFFFFFFFF00000000) | av 122 *b_p = (*b_p & 0xFFFFFFFF00000000) | bv 123 i = i + 1 124 } 125 126 println("=== nx_dot_simd_demo N=8192 i32 elements ===" as *u8) 127 128 // ---- 1. scalar ---- 129 let t0_s: i64 = nx_clock_monotonic_ns() 130 let s_scalar: i64 = scalar_dot(a_raw, b_raw, N) 131 let t1_s: i64 = nx_clock_monotonic_ns() 132 println("scalar result:" as *u8); print_i64(s_scalar); println("" as *u8) 133 println("scalar ns:" as *u8); print_i64(t1_s - t0_s); println("" as *u8) 134 135 // ---- 2. SIMD ---- 136 let t0_v: i64 = nx_clock_monotonic_ns() 137 let s_simd: i64 = simd_dot(a_raw, b_raw, N) 138 let t1_v: i64 = nx_clock_monotonic_ns() 139 println("simd result:" as *u8); print_i64(s_simd); println("" as *u8) 140 println("simd ns:" as *u8); print_i64(t1_v - t0_v); println("" as *u8) 141 if s_simd != s_scalar { println("FAIL: SIMD != scalar" as *u8); return 1 } 142 143 // ---- 3. parallel + SIMD ---- 144 var n_workers: i64 = nx_hw_worker_count() 145 if n_workers < 2 { n_workers = 2 } 146 let pool: *NxThreadPool = nx_pool_new(n_workers, 16) 147 let chunk_size: i64 = N / n_workers 148 let acc_raw: *u8 = sys_mmap(16) 149 let acc_p: *i64 = acc_raw as *i64 150 *acc_p = 0 151 _set_dot_ctx(a_raw as i64, b_raw as i64, acc_raw as i64, chunk_size) 152 153 let t0_p: i64 = nx_clock_monotonic_ns() 154 nx_parallel_for(pool, 0, n_workers, simd_dot_worker) 155 let s_par: i64 = nx_atom_load_i64(acc_p, NX_MO_SEQ_CST) 156 let t1_p: i64 = nx_clock_monotonic_ns() 157 println("par+simd result:" as *u8); print_i64(s_par); println("" as *u8) 158 println("par+simd ns:" as *u8); print_i64(t1_p - t0_p); println("" as *u8) 159 if s_par != s_scalar { println("FAIL: parallel+SIMD != scalar" as *u8); return 2 } 160 161 nx_pool_shutdown(pool) 162 println("" as *u8) 163 println("All 3 paths bit-exact. L7 x L8 composed dot product end-to-end." as *u8) 164 return 0 165}