code wiki / (root) / nx_task_graph_test.nx

nx_task_graph_test.nx source

↩ module page · 127 lines · 4552 B

1// nx_task_graph_test.nx -- DAG ordering + dependency satisfaction. 2// 3// T1: diamond DAG 4// A -> B 5// A -> C 6// B -> D 7// C -> D 8// Each node writes its tag (1..4) into a sequence log at an 9// atomic-FAA index. D must appear AFTER both B and C. 10// 11// T2: fan-out / fan-in (1 source, K fan-out, 1 sink) with hw-sized 12// parallelism inside the fan. 13 14import "nx_kernel_v2.nx" 15import "nx_log.nx" 16import "nx_atom.nx" 17import "nx_thread_pool.nx" 18import "nx_task_graph.nx" 19import "nx_hw.nx" 20 21const SEQ_LOG_CAP: i64 = 64 22 23// Shared state passed to each node via ctx (we pass the address of 24// this struct as i64). 25struct TraceState { 26 seq_idx: i64, // atomic; bumped per node finish 27 seq_log_addr: i64, // *i64 to log array 28 tag_a: i64, tag_b: i64, tag_c: i64, tag_d: i64, 29 pos_b: i64, pos_c: i64, pos_d: i64, 30} 31 32func _trace_log(state_addr: i64, tag: i64) -> i64 { 33 let s: *TraceState = state_addr as *TraceState 34 let log: *i64 = s.seq_log_addr as *i64 35 let seq_addr: *i64 = (state_addr as *i64) 36 let idx: i64 = nx_atom_faa_i64(seq_addr, 1, NX_MO_SEQ_CST) 37 log[idx] = tag 38 return idx 39} 40 41func node_A(ctx: i64) -> i64 { 42 _trace_log(ctx, 1) 43 return 0 44} 45func node_B(ctx: i64) -> i64 { 46 let pos: i64 = _trace_log(ctx, 2) 47 let s: *TraceState = ctx as *TraceState 48 s.pos_b = pos 49 return 0 50} 51func node_C(ctx: i64) -> i64 { 52 let pos: i64 = _trace_log(ctx, 3) 53 let s: *TraceState = ctx as *TraceState 54 s.pos_c = pos 55 return 0 56} 57func node_D(ctx: i64) -> i64 { 58 let pos: i64 = _trace_log(ctx, 4) 59 let s: *TraceState = ctx as *TraceState 60 s.pos_d = pos 61 return 0 62} 63 64func main() -> nx_exit { 65 var n_workers: i64 = nx_hw_worker_count() 66 if n_workers < 2 { n_workers = 2 } 67 let pool: *NxThreadPool = nx_pool_new(n_workers, 64) 68 println("=== nx_task_graph DAG smoke ===" as *u8) 69 print_i64(n_workers); println(" pool workers" as *u8) 70 71 // ---- T1: diamond DAG ---- 72 let log_raw: *u8 = sys_mmap(SEQ_LOG_CAP * 8) 73 let state_raw: *u8 = sys_mmap(128) 74 let st: *TraceState = state_raw as *TraceState 75 st.seq_idx = 0 76 st.seq_log_addr = log_raw as i64 77 st.pos_b = -1; st.pos_c = -1; st.pos_d = -1 78 79 let g: *NxTaskGraph = nx_graph_new(pool, 16) 80 let a_idx: i64 = nx_graph_add_node(g, node_A, state_raw as i64) 81 let b_idx: i64 = nx_graph_add_node(g, node_B, state_raw as i64) 82 let c_idx: i64 = nx_graph_add_node(g, node_C, state_raw as i64) 83 let d_idx: i64 = nx_graph_add_node(g, node_D, state_raw as i64) 84 nx_graph_add_edge(g, a_idx, b_idx) 85 nx_graph_add_edge(g, a_idx, c_idx) 86 nx_graph_add_edge(g, b_idx, d_idx) 87 nx_graph_add_edge(g, c_idx, d_idx) 88 89 if nx_graph_run(g) != 0 { println("FAIL T1 run timeout" as *u8); return 1 } 90 if nx_graph_n_completed(g) != 4 { println("FAIL T1 completed_count" as *u8); return 2 } 91 92 // Order checks: D must come after B AND C. A must be at pos 0. 93 let log: *i64 = log_raw as *i64 94 if log[0] != 1 { println("FAIL T1: A not first" as *u8); return 3 } 95 if st.pos_d < st.pos_b { println("FAIL T1: D ran before B" as *u8); return 4 } 96 if st.pos_d < st.pos_c { println("FAIL T1: D ran before C" as *u8); return 5 } 97 if st.pos_d != 3 { println("FAIL T1: D not last" as *u8); return 6 } 98 println("PASS T1: diamond DAG ordering correct (A->{B,C}->D)" as *u8) 99 100 // ---- T2: fan-out source -> 8 leaves -> sink ---- 101 let log2_raw: *u8 = sys_mmap(SEQ_LOG_CAP * 8) 102 let state2_raw: *u8 = sys_mmap(128) 103 let st2: *TraceState = state2_raw as *TraceState 104 st2.seq_idx = 0 105 st2.seq_log_addr = log2_raw as i64 106 st2.pos_b = -1; st2.pos_c = -1; st2.pos_d = -1 107 108 let g2: *NxTaskGraph = nx_graph_new(pool, 16) 109 let src: i64 = nx_graph_add_node(g2, node_A, state2_raw as i64) 110 let sink: i64 = nx_graph_add_node(g2, node_D, state2_raw as i64) 111 var k: i64 = 0 112 while k < 8 { 113 let mid: i64 = nx_graph_add_node(g2, node_B, state2_raw as i64) 114 nx_graph_add_edge(g2, src, mid) 115 nx_graph_add_edge(g2, mid, sink) 116 k = k + 1 117 } 118 if nx_graph_run(g2) != 0 { println("FAIL T2 run timeout" as *u8); return 7 } 119 if nx_graph_n_completed(g2) != 10 { println("FAIL T2 completed_count" as *u8); return 8 } 120 if st2.pos_d < 9 { println("FAIL T2 sink ran early" as *u8); return 9 } 121 println("PASS T2: fan-out (1 -> 8 -> 1) -- sink ran last" as *u8) 122 123 nx_pool_shutdown(pool) 124 println("" as *u8) 125 println("Task graph proven: dependency edges enforced under MIMD scheduling." as *u8) 126 return 0 127}