nx_eqsat_test.nx source
↩ module page · 90 lines · 4997 B
1// nx_eqsat_test.nx -- proof-of-life + correctness GATE for the equality-
2// saturation engine (the GENERATOR organ of the Sovereign Invention Engine).
3//
4// Case 1 -- EXTRACTOR correctness. (add (mul x x) 0): add_zero unions it with
5// (mul x x); the add-class stays canonical (shallow cost 1 < mul's 3) with
6// best_node = the add node, so the V1 CACHED extractor returns
7// (add (mul x x) 0) [cost 4] -- a wasted +0. The REAL bottom-up extractor
8// (nx_eqsat_recompute_best) returns (mul x x) [cost 3]. Known: before_op=ADD(2),
9// cached_op=ADD(2) [the bug], real_op=MUL(4), real_cost=3.
10//
11// Case 2 -- a REAL optimization WIN. (mul x 8): the strength-reduction rule
12// (mul x 2^k)->(shl x k) fires, and the real extractor picks (shl x 3)
13// [cost 1] over (mul x 8) [cost 3] -- a genuine 3->1 win. The idempotent rule
14// lets saturation CONVERGE (NX_EQSAT_SATURATED). Known: red_op=SHL(11),
15// red_cost=1.
16//
17// Full known answer (FAIL LOUD): "2 2 4 3 11 1 ".
18
19import "nx_eqsat.nx"
20
21func _emit_num(v: i64) -> i64 {
22 let b: *u8 = sys_mmap(28); var n: i64 = v; if n < 0 { n = 0 - n }
23 let t2: *u8 = sys_mmap(28); var t: i64 = 0
24 if n == 0 { t2[0] = 48; t = 1 }
25 while n > 0 { t2[t] = 48 + (n % 10); n = n / 10; t = t + 1 }
26 var i: i64 = 0; while i < t { b[i] = t2[t - 1 - i]; i = i + 1 }
27 b[t] = 32; sys_write(1, b, t + 1); return 0
28}
29func _nl() -> i64 { let z: *u8 = sys_mmap(2); z[0] = 10; sys_write(1, z, 1); return 0 }
30
31func main() -> i64 {
32 // ---- Case 1: extractor discriminator -- (add (mul x x) 0) ----
33 let nodes: *NxENode = sys_mmap(256 * 128) as *NxENode
34 let classes: *NxEClass = sys_mmap(256 * 64) as *NxEClass
35 let g: *NxEGraph = sys_mmap(256) as *NxEGraph
36 if nx_eqsat_init(g, nodes, 256, classes, 256) != NX_EQSAT_OK { sys_exit(10); return 10 }
37 let x: i64 = nx_eqsat_add_var(g, 0)
38 let mul_xx: i64 = nx_eqsat_add_binary(g, NX_EQ_OP_MUL, x, x)
39 let zero: i64 = nx_eqsat_add_const(g, 0)
40 let top: i64 = nx_eqsat_add_binary(g, NX_EQ_OP_ADD, mul_xx, zero)
41 let before_op: i64 = g.nodes[nx_eqsat_extract_best_node(g, top)].op
42 let src1: i64 = nx_eqsat_saturate(g, 16)
43 let cached_op: i64 = g.nodes[nx_eqsat_extract_best_node(g, top)].op
44 if nx_eqsat_recompute_best(g) != NX_EQSAT_OK { sys_exit(11); return 11 }
45 let real_op: i64 = g.nodes[nx_eqsat_extract_best_node(g, top)].op
46 let real_cost: i64 = nx_eqsat_best_cost(g, top)
47 _emit_num(before_op); _emit_num(cached_op); _emit_num(real_op); _emit_num(real_cost)
48
49 // ---- Case 2: real strength-reduction win -- (mul x 8) -> (shl x 3) ----
50 let nodes2: *NxENode = sys_mmap(256 * 128) as *NxENode
51 let classes2: *NxEClass = sys_mmap(256 * 64) as *NxEClass
52 let g2: *NxEGraph = sys_mmap(256) as *NxEGraph
53 if nx_eqsat_init(g2, nodes2, 256, classes2, 256) != NX_EQSAT_OK { sys_exit(12); return 12 }
54 let x2: i64 = nx_eqsat_add_var(g2, 0)
55 let eight: i64 = nx_eqsat_add_const(g2, 8)
56 let m2: i64 = nx_eqsat_add_binary(g2, NX_EQ_OP_MUL, x2, eight)
57 let src2: i64 = nx_eqsat_saturate(g2, 16)
58 if nx_eqsat_recompute_best(g2) != NX_EQSAT_OK { sys_exit(13); return 13 }
59 let red_op: i64 = g2.nodes[nx_eqsat_extract_best_node(g2, m2)].op
60 let red_cost: i64 = nx_eqsat_best_cost(g2, m2)
61 _emit_num(red_op); _emit_num(red_cost)
62
63 // ---- Case 3: program-emitting extractor -- emit (shl x 3) as a real DAG ----
64 let out: *NxEmitNode = sys_mmap(64 * 40) as *NxEmitNode
65 let cnt: *i64 = sys_mmap(8) as *i64
66 cnt[0] = 0
67 let root: i64 = nx_eqsat_emit(g2, m2, out, 64, cnt)
68 if root < 0 { sys_exit(14); return 14 }
69 let emit_count: i64 = cnt[0]
70 let root_op: i64 = out[root].op
71 let shift_amt: i64 = out[out[root].c1].payload
72 let base_op: i64 = out[out[root].c0].op
73 _emit_num(emit_count); _emit_num(shift_amt)
74 _nl()
75
76 // ---- assertions (FAIL LOUD) ----
77 if before_op != NX_EQ_OP_ADD { sys_exit(1); return 1 } // started as ADD
78 if real_op != NX_EQ_OP_MUL { sys_exit(2); return 2 } // extractor dropped the +0
79 if real_cost != 3 { sys_exit(3); return 3 } // (mul x x)=3, not (add..0)=4
80 if nx_eqsat_find(g, top) != nx_eqsat_find(g, mul_xx) { sys_exit(4); return 4 } // union fired
81 if cached_op != NX_EQ_OP_ADD { sys_exit(5); return 5 } // documents the V1 bug fixed
82 if red_op != NX_EQ_OP_SHL { sys_exit(6); return 6 } // strength reduction fired
83 if red_cost != 1 { sys_exit(7); return 7 } // genuine 3 -> 1 cost win
84 if src2 != NX_EQSAT_SATURATED { sys_exit(8); return 8 } // idempotent rule converged
85 if emit_count != 3 { sys_exit(20); return 20 } // emitted exactly (shl x 3): x, 3, shl
86 if root_op != NX_EQ_OP_SHL { sys_exit(21); return 21 } // root of the emitted program is SHL
87 if shift_amt != 3 { sys_exit(22); return 22 } // shift amount = log2(8) = 3
88 if base_op != NX_EQ_OP_VAR { sys_exit(23); return 23 } // shifted operand is the var x
89 sys_exit(0); return 0
90}