code wiki / (root) / nx_parallel_test.nx

nx_parallel_test.nx source

↩ module page · 99 lines · 3720 B

1// nx_parallel_test.nx -- larger-scale parallel primitive stress. 2// 3// T1: parallel_map of N elements where every input is the counter 4// address; fn FAA-bumps the counter. Functionally equivalent 5// to parallel_for(0, N, bump) but avoids the pre-existing 6// static-i64-from-main codegen bug documented in 7// [[project-nx-threads-bedrock-2026-05-15]] pitfalls. 8// T2: parallel_map: square 4096 elements, verify each. 9// T3: parallel_reduce: sum 1..1000 == 500500. 10 11import "nx_kernel_v2.nx" 12import "nx_log.nx" 13import "nx_atom.nx" 14import "nx_thread_pool.nx" 15import "nx_parallel.nx" 16import "nx_hw.nx" 17 18const T1_ITERS: i64 = 4000 19const MAP_LEN: i64 = 4096 20const REDUCE_LEN: i64 = 1000 // sum 1..1000 = 1000*1001/2 = 500500 21 22// Per-element fn that treats the value as a pointer address and 23// FAA-bumps it. Returns 0 (we ignore the map output). 24func bump_via_addr(addr_as_i64: i64) -> i64 { 25 let p: *i64 = addr_as_i64 as *i64 26 nx_atom_faa_i64(p, 1, NX_MO_SEQ_CST) 27 return 0 28} 29 30func square(x: i64) -> i64 { return x * x } 31func iadd(a: i64, b: i64) -> i64 { return a + b } 32 33func main() -> nx_exit { 34 var n_workers: i64 = nx_hw_worker_count() 35 if n_workers < 2 { n_workers = 2 } // ensure contention even on 1-cpu qemu 36 let pool: *NxThreadPool = nx_pool_new(n_workers, 64) 37 38 println("=== nx_parallel stress smoke ===" as *u8) 39 println("Pool workers (max(hw, 2)):" as *u8); print_i64(n_workers); println("" as *u8) 40 41 // ---- T1: parallel_map-as-parallel_for FAA stress ---- 42 let counter_raw: *u8 = sys_mmap(16) 43 let counter: *i64 = counter_raw as *i64 44 *counter = 0 45 let counter_addr: i64 = counter as i64 46 47 let in_raw_1: *u8 = sys_mmap(T1_ITERS * 8) 48 let out_raw_1: *u8 = sys_mmap(T1_ITERS * 8) 49 let in_arr_1: *i64 = in_raw_1 as *i64 50 let out_arr_1: *i64 = out_raw_1 as *i64 51 var i: i64 = 0 52 while i < T1_ITERS { in_arr_1[i] = counter_addr; i = i + 1 } 53 54 if nx_parallel_map_i64(pool, in_arr_1, out_arr_1, T1_ITERS, bump_via_addr) != 0 { 55 println("FAIL T1 parallel_map return" as *u8); return 1 56 } 57 let v1: i64 = nx_atom_load_i64(counter, NX_MO_SEQ_CST) 58 if v1 != T1_ITERS { 59 println("FAIL T1: counter != T1_ITERS" as *u8); return 2 60 } 61 println("PASS T1: parallel_map 4000 elements, counter == 4000 (no lost FAAs)" as *u8) 62 63 // ---- T2: parallel_map square 4096 elements ---- 64 let in_raw_2: *u8 = sys_mmap(MAP_LEN * 8) 65 let out_raw_2: *u8 = sys_mmap(MAP_LEN * 8) 66 let in_arr_2: *i64 = in_raw_2 as *i64 67 let out_arr_2: *i64 = out_raw_2 as *i64 68 var j: i64 = 0 69 while j < MAP_LEN { in_arr_2[j] = j; j = j + 1 } 70 if nx_parallel_map_i64(pool, in_arr_2, out_arr_2, MAP_LEN, square) != 0 { 71 println("FAIL T2 parallel_map return" as *u8); return 3 72 } 73 var jj: i64 = 0 74 var fail2: i64 = 0 75 while jj < MAP_LEN { 76 if out_arr_2[jj] != jj * jj { fail2 = jj + 1 } 77 jj = jj + 1 78 } 79 if fail2 != 0 { 80 println("FAIL T2: out[i] != i*i" as *u8); return 4 81 } 82 println("PASS T2: parallel_map 4096 squares all correct" as *u8) 83 84 // ---- T3: parallel_reduce sum 1..1000 ---- 85 let r_raw: *u8 = sys_mmap(REDUCE_LEN * 8) 86 let r_arr: *i64 = r_raw as *i64 87 var k: i64 = 0 88 while k < REDUCE_LEN { r_arr[k] = k + 1; k = k + 1 } 89 let sum: i64 = nx_parallel_reduce_i64(pool, r_arr, REDUCE_LEN, 0, iadd) 90 if sum != 500500 { 91 println("FAIL T3: sum != 500500" as *u8); print_i64(sum); println("" as *u8); return 5 92 } 93 println("PASS T3: parallel_reduce sum 1..1000 == 500500" as *u8) 94 95 nx_pool_shutdown(pool) 96 println("" as *u8) 97 println("All parallel primitives green; pool drained cleanly." as *u8) 98 return 0 99}