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}