nx_eqsat_vs_gcc_battery_test.nx source
↩ module page · 178 lines · 8200 B
1// nx_eqsat_vs_gcc_battery_test.nx -- FAIR apples-to-apples op-count battery
2// for the e-graph optimizer vs gcc -O3 (incumbent oracle, same box).
3//
4// METRIC (stated explicitly): ARITHMETIC OPERATION COUNT of the optimized
5// straight-line computation. We count DISTINCT arithmetic nodes in the
6// emitted DAG (op != CONST, op != VAR), i.e. each shared subexpression once
7// -- the SAME thing one counts in gcc's asm (add/sub/imul/lea-scale/shl/
8// shr/sar/and/or/xor/not/neg in the function body, excluding prologue/
9// epilogue/reg-moves/ret). This is NOT the recursive class_cost (which
10// double-counts shared children and would UNFAIRLY inflate our number).
11//
12// NON-CHERRY-PICKED: a fixed battery declared up front; every kernel reported
13// whatever the verdict (no dropping losers). Each kernel expresses the SAME
14// computation we hand to gcc in battery.c.
15//
16// Output per kernel: "<id> <emit_total> <arith_ops> <cost> <root_op> <sat>"
17// arith_ops is the FAIR metric. cost is class_cost (context only). sat is the
18// saturate status (4=SATURATED/fixpoint, 5=hit-budget; extraction is sound
19// either way -- costs only decrease, so budget exhaustion is not an error).
20
21import "nx_eqsat.nx"
22
23func _emit_num(v: i64) -> i64 {
24 let b: *u8 = sys_mmap(28); var n: i64 = v; if n < 0 { n = 0 - n }
25 let t2: *u8 = sys_mmap(28); var t: i64 = 0
26 if n == 0 { t2[0] = 48; t = 1 }
27 while n > 0 { t2[t] = 48 + (n % 10); n = n / 10; t = t + 1 }
28 var i: i64 = 0; while i < t { b[i] = t2[t - 1 - i]; i = i + 1 }
29 b[t] = 32; sys_write(1, b, t + 1); return 0
30}
31func _nl() -> i64 { let z: *u8 = sys_mmap(2); z[0] = 10; sys_write(1, z, 1); return 0 }
32
33// Structural equality of two emitted subtrees (recursive; small DAGs).
34func _emit_eq(out: *NxEmitNode, i: i64, j: i64) -> i64 {
35 if i == j { return 1 }
36 if out[i].op != out[j].op { return 0 }
37 if out[i].payload != out[j].payload { return 0 }
38 let ci0: i64 = out[i].c0; let cj0: i64 = out[j].c0
39 let ci1: i64 = out[i].c1; let cj1: i64 = out[j].c1
40 let ci2: i64 = out[i].c2; let cj2: i64 = out[j].c2
41 if ci0 < 0 { if cj0 >= 0 { return 0 } } else { if cj0 < 0 { return 0 } else { if _emit_eq(out, ci0, cj0) == 0 { return 0 } } }
42 if ci1 < 0 { if cj1 >= 0 { return 0 } } else { if cj1 < 0 { return 0 } else { if _emit_eq(out, ci1, cj1) == 0 { return 0 } } }
43 if ci2 < 0 { if cj2 >= 0 { return 0 } } else { if cj2 < 0 { return 0 } else { if _emit_eq(out, ci2, cj2) == 0 { return 0 } } }
44 return 1
45}
46
47// Count DISTINCT arithmetic ops (op != CONST/VAR), deduping structurally-equal
48// shared subexpressions -- the FAIR count matching gcc's CSE (a shared (a+b)
49// counts ONCE). nx_eqsat_emit produces a post-order TREE (re-emits shared
50// children), so we must dedupe to compare apples-to-apples.
51func _arith_ops(out: *NxEmitNode, cnt: i64) -> i64 {
52 var k: i64 = 0; var a: i64 = 0
53 while k < cnt {
54 let op: i64 = out[k].op
55 if op != NX_EQ_OP_CONST {
56 if op != NX_EQ_OP_VAR {
57 // count k only if no EARLIER node is structurally identical
58 var seen: i64 = 0; var j: i64 = 0
59 while j < k {
60 if out[j].op != NX_EQ_OP_CONST {
61 if out[j].op != NX_EQ_OP_VAR {
62 if _emit_eq(out, j, k) == 1 { seen = 1; j = k }
63 }
64 }
65 j = j + 1
66 }
67 if seen == 0 { a = a + 1 }
68 }
69 }
70 k = k + 1
71 }
72 return a
73}
74
75// Saturate + recompute + emit + report one kernel rooted at `root`.
76func _report(id: i64, g: *NxEGraph, root: i64) -> i64 {
77 let s: i64 = nx_eqsat_saturate(g, 64)
78 // Budget exhaustion (-NX_EQSAT_STEP_BUDGET) is NOT an error: the e-graph is
79 // valid and recompute_best still extracts the sound minimum-so-far. Only a
80 // DIFFERENT negative (real failure) is fatal.
81 var sat: i64 = s
82 if s < 0 {
83 if s != (0 - NX_EQSAT_STEP_BUDGET) { sys_exit(40 + id); return 40 + id }
84 sat = NX_EQSAT_STEP_BUDGET
85 }
86 if nx_eqsat_recompute_best(g) != NX_EQSAT_OK { sys_exit(60 + id); return 60 + id }
87 let cost: i64 = nx_eqsat_best_cost(g, root)
88 let out: *NxEmitNode = sys_mmap(256 * 40) as *NxEmitNode
89 let cnt: *i64 = sys_mmap(8) as *i64
90 cnt[0] = 0
91 let r: i64 = nx_eqsat_emit(g, root, out, 256, cnt)
92 if r < 0 { sys_exit(80 + id); return 80 + id }
93 let total: i64 = cnt[0]
94 let arith: i64 = _arith_ops(out, total)
95 let rop: i64 = out[r].op
96 _emit_num(id); _emit_num(total); _emit_num(arith); _emit_num(cost); _emit_num(rop); _emit_num(sat); _nl()
97 return arith
98}
99
100func _fresh(g_out: *i64) -> *NxEGraph {
101 let nodes: *NxENode = sys_mmap(512 * 128) as *NxENode
102 let classes: *NxEClass = sys_mmap(512 * 64) as *NxEClass
103 let g: *NxEGraph = sys_mmap(256) as *NxEGraph
104 if nx_eqsat_init(g, nodes, 512, classes, 512) != NX_EQSAT_OK { sys_exit(20); }
105 // Enable the const-fold e-class analysis the engine already ships (opt-in).
106 // This is a REAL capability gcc also has -- fair to turn on, not a cheat.
107 if nx_eqsat_enable_constfold(g) < 0 { sys_exit(21); }
108 return g
109}
110
111func main() -> i64 {
112 // K1: y = x*8 (strength reduction -> shl x 3)
113 let g1: *NxEGraph = _fresh(0 as *i64)
114 let x1: i64 = nx_eqsat_add_var(g1, 0)
115 let c8: i64 = nx_eqsat_add_const(g1, 8)
116 let k1: i64 = nx_eqsat_add_binary(g1, NX_EQ_OP_MUL, x1, c8)
117 let a1: i64 = _report(1, g1, k1)
118
119 // K2: y = (a+b)*(a+b) (CSE -> one add, one mul)
120 let g2: *NxEGraph = _fresh(0 as *i64)
121 let a_2: i64 = nx_eqsat_add_var(g2, 0)
122 let b_2: i64 = nx_eqsat_add_var(g2, 1)
123 let s2: i64 = nx_eqsat_add_binary(g2, NX_EQ_OP_ADD, a_2, b_2)
124 let k2: i64 = nx_eqsat_add_binary(g2, NX_EQ_OP_MUL, s2, s2)
125 let a2: i64 = _report(2, g2, k2)
126
127 // K3: y = x*2 + x*2 (algebraic + strength -> shl x 2)
128 let g3: *NxEGraph = _fresh(0 as *i64)
129 let x3: i64 = nx_eqsat_add_var(g3, 0)
130 let c2: i64 = nx_eqsat_add_const(g3, 2)
131 let m3a: i64 = nx_eqsat_add_binary(g3, NX_EQ_OP_MUL, x3, c2)
132 let m3b: i64 = nx_eqsat_add_binary(g3, NX_EQ_OP_MUL, x3, c2)
133 let k3: i64 = nx_eqsat_add_binary(g3, NX_EQ_OP_ADD, m3a, m3b)
134 let a3: i64 = _report(3, g3, k3)
135
136 // K4: y = (x + 0) * 1 (identity elimination -> x, 0 arith ops)
137 let g4: *NxEGraph = _fresh(0 as *i64)
138 let x4: i64 = nx_eqsat_add_var(g4, 0)
139 let z4: i64 = nx_eqsat_add_const(g4, 0)
140 let o4: i64 = nx_eqsat_add_const(g4, 1)
141 let p4: i64 = nx_eqsat_add_binary(g4, NX_EQ_OP_ADD, x4, z4)
142 let k4: i64 = nx_eqsat_add_binary(g4, NX_EQ_OP_MUL, p4, o4)
143 let a4: i64 = _report(4, g4, k4)
144
145 // K5: y = (x - x) + (a & a) (sub_self -> 0, and_self -> a; -> a)
146 let g5: *NxEGraph = _fresh(0 as *i64)
147 let x5: i64 = nx_eqsat_add_var(g5, 0)
148 let a_5: i64 = nx_eqsat_add_var(g5, 1)
149 let ss5: i64 = nx_eqsat_add_binary(g5, NX_EQ_OP_SUB, x5, x5)
150 let aa5: i64 = nx_eqsat_add_binary(g5, NX_EQ_OP_AND, a_5, a_5)
151 let k5: i64 = nx_eqsat_add_binary(g5, NX_EQ_OP_ADD, ss5, aa5)
152 let a5: i64 = _report(5, g5, k5)
153
154 // K6: y = (x*4) * 8 (const-fold mul chain -> shl x 5)
155 let g6: *NxEGraph = _fresh(0 as *i64)
156 let x6: i64 = nx_eqsat_add_var(g6, 0)
157 let c4: i64 = nx_eqsat_add_const(g6, 4)
158 let c8b: i64 = nx_eqsat_add_const(g6, 8)
159 let m6: i64 = nx_eqsat_add_binary(g6, NX_EQ_OP_MUL, x6, c4)
160 let k6: i64 = nx_eqsat_add_binary(g6, NX_EQ_OP_MUL, m6, c8b)
161 let a6: i64 = _report(6, g6, k6)
162
163 // K7: y = a + (a + a) (add_self folding: a+a -> shl a 1)
164 let g7: *NxEGraph = _fresh(0 as *i64)
165 let a_7: i64 = nx_eqsat_add_var(g7, 0)
166 let inner7: i64 = nx_eqsat_add_binary(g7, NX_EQ_OP_ADD, a_7, a_7)
167 let k7: i64 = nx_eqsat_add_binary(g7, NX_EQ_OP_ADD, a_7, inner7)
168 let a7: i64 = _report(7, g7, k7)
169
170 // K8: y = 6 + 7 (pure const-fold -> 13, 0 arith ops)
171 let g8: *NxEGraph = _fresh(0 as *i64)
172 let c6: i64 = nx_eqsat_add_const(g8, 6)
173 let c7: i64 = nx_eqsat_add_const(g8, 7)
174 let k8: i64 = nx_eqsat_add_binary(g8, NX_EQ_OP_ADD, c6, c7)
175 let a8: i64 = _report(8, g8, k8)
176
177 return 0
178}