code wiki / (root) / nx_q4k_matmul_mt_gate.nx

nx_q4k_matmul_mt_gate.nx source

↩ module page · 240 lines · 7861 B

1// nx_q4k_matmul_mt_gate.nx -- adversarial gate for the multi-threaded 2// Q4_K matmul (nx_f32_q4k_matmul_mt / _pool) against the serial oracle 3// nx_f32_q4k_matmul. This is the LLM forward's hot path: the organ 4// recorded "pool (fork-once) is the fix" when fork-per-matmul proved 5// net-slower -- the thread pool IS that fix, landed 2026-07-07. 6// 7// Bit-exactness is a hard requirement and holds by construction 8// (column banding: every C[i,j] computed wholly inside one band, 9// identical accumulation order) -- the gate verifies it with full- 10// buffer i64 compares on POISONED outputs so unwritten columns can 11// never pass. Weight bytes are DUMMY (LCG): nx_q4k_to_f32 decodes 12// any bytes deterministically, so exactness + timing are layout-real 13// without a model file (nx_q4k_matmul_rate precedent). 14// 15// Checks (10): 16// 1 serial oracle OK (m=3, k=512, n=37 -- prime n for band edges) 17// 2..6 MT bit-exact for nworkers in {1,2,3,0=auto,1000=clamp-to-37} 18// 7 m=1 single-token decode shape bit-exact (auto workers) 19// 8..9 caller-owned pool reused across two calls (delta-wait), exact 20// 10 speedup, decode shape m=1 k=1024 n=4864 x6 reps: serial vs 21// POOL path (spawn-once, submit-per-call = the forward's 22// pattern); exact AND wall-time speedup >= floor 23// 24// genealogy_id: substrate_f32_q4k_matmul_v2_parallel + nx_conv2d_mt_gate 25// lineage_id: q4k_matmul_mt_gate_v1 26 27import "nx_f32_q4k_matmul.nx" 28import "nx_fmt.nx" 29 30const QG_M: i64 = 3 31const QG_K: i64 = 512 32const QG_N: i64 = 37 33 34const QB_K: i64 = 1024 35const QB_N: i64 = 4864 36const QB_REPS: i64 = 6 37 38// Conservative floor (2.0x) for a 16-worker host; measured value is 39// printed above the verdict for the honest record. 40const QG_SPEEDUP_FLOOR_X100: i64 = 200 41 42func q_lcg(s: i64) -> i64 { 43 var x: i64 = s * 1103515245 + 12345 44 x = x & 2147483647 45 return x 46} 47 48// A-matrix fill: boxed f32 of small ints in [-4, 4]. 49func q_fill_a(p: *i64, count: i64, seed: i64) -> i64 { 50 var s: i64 = seed 51 var i: i64 = 0 52 while i < count { 53 s = q_lcg(s) 54 let v: i64 = (s % 9) - 4 55 p[i] = nx_i32_to_f32(v) 56 i = i + 1 57 } 58 return 0 59} 60 61// Dummy Q4_K weight bytes (any bytes decode deterministically). 62func q_fill_b(p: *u8, count: i64, seed: i64) -> i64 { 63 var s: i64 = seed 64 var i: i64 = 0 65 while i < count { 66 s = q_lcg(s) 67 p[i] = (s & 255) as u8 68 i = i + 1 69 } 70 return 0 71} 72 73func q_poison(p: *i64, count: i64) -> i64 { 74 let pv: i64 = 0 - 777777 75 var i: i64 = 0 76 while i < count { 77 p[i] = pv 78 i = i + 1 79 } 80 return 0 81} 82 83func q_same(a: *i64, b: *i64, count: i64) -> i64 { 84 var i: i64 = 0 85 while i < count { 86 if a[i] != b[i] { return 0 } 87 i = i + 1 88 } 89 return 1 90} 91 92func q_nl() -> i64 { 93 fmt_puts("\n" as *u8) 94 return 0 95} 96 97func main() -> i64 { 98 // ---- small shape buffers ---- 99 let a_n: i64 = QG_M * QG_K 100 let b_n: i64 = QG_N * (QG_K / 256) * 144 101 let c_n: i64 = QG_M * QG_N 102 let A: *i64 = sys_mmap(a_n * 8) as *i64 103 let B: *u8 = sys_mmap(b_n) 104 let Cref: *i64 = sys_mmap(c_n * 8) as *i64 105 let Cout: *i64 = sys_mmap(c_n * 8) as *i64 106 q_fill_a(A, a_n, 20260707) 107 q_fill_b(B, b_n, 977) 108 109 var pass: i64 = 0 110 111 // ---- 1: serial oracle ---- 112 let v1: nx_int = nx_f32_q4k_matmul(A, B, 0, Cref, QG_M, QG_K, QG_N) 113 if v1 != NX_FQ4M_OK { 114 fmt_puts("Q4KMT 1 SERIAL FAIL v="); fmt_putn(v1); q_nl() 115 return 11 116 } 117 fmt_puts("Q4KMT 1 SERIAL OK"); q_nl() 118 pass = pass + 1 119 120 // ---- 2..6: MT exactness across worker counts ---- 121 let counts: *i64 = sys_mmap(5 * 8) as *i64 122 counts[0] = 1; counts[1] = 2; counts[2] = 3; counts[3] = 0; counts[4] = 1000 123 var t: i64 = 0 124 while t < 5 { 125 q_poison(Cout, c_n) 126 let nwk: i64 = counts[t] 127 let vk: nx_int = nx_f32_q4k_matmul_mt(A, B, 0, Cout, QG_M, QG_K, QG_N, nwk) 128 var ok: i64 = 0 129 if vk == NX_FQ4M_OK { ok = q_same(Cout, Cref, c_n) } 130 if ok != 1 { 131 fmt_puts("Q4KMT EXACT FAIL nw="); fmt_putn(nwk); q_nl() 132 return 12 + t 133 } 134 fmt_puts("Q4KMT EXACT OK nw="); fmt_putn(nwk); q_nl() 135 pass = pass + 1 136 t = t + 1 137 } 138 139 // ---- 7: m=1 single-token decode shape ---- 140 let c1_n: i64 = QG_N 141 let C1ref: *i64 = sys_mmap(c1_n * 8) as *i64 142 let C1out: *i64 = sys_mmap(c1_n * 8) as *i64 143 let v7s: nx_int = nx_f32_q4k_matmul(A, B, 0, C1ref, 1, QG_K, QG_N) 144 q_poison(C1out, c1_n) 145 let v7m: nx_int = nx_f32_q4k_matmul_mt(A, B, 0, C1out, 1, QG_K, QG_N, 0) 146 var ok7: i64 = 0 147 if v7s == NX_FQ4M_OK { if v7m == NX_FQ4M_OK { ok7 = q_same(C1out, C1ref, c1_n) } } 148 if ok7 != 1 { 149 fmt_puts("Q4KMT 7 M1 FAIL"); q_nl() 150 return 17 151 } 152 fmt_puts("Q4KMT 7 M1 OK"); q_nl() 153 pass = pass + 1 154 155 // ---- 8..9: caller-owned pool reused across two calls ---- 156 let pool4: *NxThreadPool = nx_pool_new(4, 0) 157 q_poison(Cout, c_n) 158 let v8: nx_int = nx_f32_q4k_matmul_pool(pool4, A, B, 0, Cout, QG_M, QG_K, QG_N) 159 var ok8: i64 = 0 160 if v8 == NX_FQ4M_OK { ok8 = q_same(Cout, Cref, c_n) } 161 if ok8 != 1 { 162 fmt_puts("Q4KMT 8 POOL1 FAIL"); q_nl() 163 return 18 164 } 165 fmt_puts("Q4KMT 8 POOL1 OK"); q_nl() 166 pass = pass + 1 167 168 q_fill_a(A, a_n, 555001) 169 let v9s: nx_int = nx_f32_q4k_matmul(A, B, 0, Cref, QG_M, QG_K, QG_N) 170 if v9s != NX_FQ4M_OK { return 19 } 171 q_poison(Cout, c_n) 172 let v9: nx_int = nx_f32_q4k_matmul_pool(pool4, A, B, 0, Cout, QG_M, QG_K, QG_N) 173 var ok9: i64 = 0 174 if v9 == NX_FQ4M_OK { ok9 = q_same(Cout, Cref, c_n) } 175 nx_pool_shutdown(pool4) 176 if ok9 != 1 { 177 fmt_puts("Q4KMT 9 POOL2 FAIL"); q_nl() 178 return 19 179 } 180 fmt_puts("Q4KMT 9 POOL2 OK"); q_nl() 181 pass = pass + 1 182 183 // ---- 10: decode-shape speedup, serial vs pool (forward pattern) ---- 184 let ba_n: i64 = QB_K 185 let bb_n: i64 = QB_N * (QB_K / 256) * 144 186 let bc_n: i64 = QB_N 187 let BA: *i64 = sys_mmap(ba_n * 8) as *i64 188 let BB: *u8 = sys_mmap(bb_n) 189 let BCref: *i64 = sys_mmap(bc_n * 8) as *i64 190 let BCout: *i64 = sys_mmap(bc_n * 8) as *i64 191 q_fill_a(BA, ba_n, 6767) 192 q_fill_b(BB, bb_n, 8181) 193 194 let t0: i64 = sys_now_us() 195 var r0: i64 = 0 196 while r0 < QB_REPS { 197 let vs: nx_int = nx_f32_q4k_matmul(BA, BB, 0, BCref, 1, QB_K, QB_N) 198 if vs != NX_FQ4M_OK { return 21 } 199 r0 = r0 + 1 200 } 201 let serial_us: i64 = sys_now_us() - t0 202 203 let poolw: *NxThreadPool = nx_pool_new(0, 0) 204 q_poison(BCout, bc_n) 205 let t1: i64 = sys_now_us() 206 var r1: i64 = 0 207 while r1 < QB_REPS { 208 let vp: nx_int = nx_f32_q4k_matmul_pool(poolw, BA, BB, 0, BCout, 1, QB_K, QB_N) 209 if vp != NX_FQ4M_OK { return 22 } 210 r1 = r1 + 1 211 } 212 let mt_us: i64 = sys_now_us() - t1 213 let nwlive: i64 = poolw.n_workers 214 nx_pool_shutdown(poolw) 215 216 let okb: i64 = q_same(BCout, BCref, bc_n) 217 if okb != 1 { 218 fmt_puts("Q4KMT 10 BIG-EXACT FAIL"); q_nl() 219 return 23 220 } 221 222 fmt_puts("pool_workers="); fmt_putn(nwlive); q_nl() 223 fmt_puts("serial_us="); fmt_putn(serial_us); q_nl() 224 fmt_puts("mt_us="); fmt_putn(mt_us); q_nl() 225 let macs: i64 = QB_K * QB_N * QB_REPS 226 if serial_us > 0 { fmt_puts("serial_mflops="); fmt_putn(2 * macs / serial_us); q_nl() } 227 if mt_us > 0 { fmt_puts("mt_mflops="); fmt_putn(2 * macs / mt_us); q_nl() } 228 var sx100: i64 = 0 229 if mt_us > 0 { sx100 = serial_us * 100 / mt_us } 230 fmt_puts("speedup_x100="); fmt_putn(sx100); q_nl() 231 if sx100 < QG_SPEEDUP_FLOOR_X100 { 232 fmt_puts("Q4KMT 10 SPEEDUP FAIL"); q_nl() 233 return 24 234 } 235 fmt_puts("Q4KMT 10 SPEEDUP OK"); q_nl() 236 pass = pass + 1 237 238 fmt_puts("Q4K_MATMUL_MT_GATE "); fmt_putn(pass); fmt_puts("/10 GREEN"); q_nl() 239 return 0 240}