code wiki / (root) / nx_eqsat_vs_gcc_battery_test.nx

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}