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}