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}