nx_gen_blockrun.nx source
↩ module page · 330 lines · 16798 B
1// nx_gen_blockrun.nx -- run a WHOLE DiT block sovereignly, end to end, and grade the output.
2//
3// Takes ONLY the block input and the model weights. Every intermediate is computed by this organ
4// and fed to the next stage; no oracle tensor is read except the final reference.
5//
6// WHY THIS EXISTS SEPARATELY FROM nx_gen_blockbench
7// The bench proves each op matches the oracle WHEN HANDED THE ORACLE'S INPUTS. That is necessary
8// and not sufficient: it says nothing about whether 17 chained ops stay on the rails, because
9// every stage there starts from a clean reference. ★ VERIFIED IS NOT ASSEMBLED. Error compounding
10// across a chain is a different question from per-op fidelity, and only this organ asks it.
11//
12// Usage: nx_gen_blockrun <model_id> [n_workers]
13//
14// All activations are kept as PACKED 4-byte f32, not i64-boxed floats, because that is the layout
15// __f32_i8dot32a consumes -- the quantized dot needs 32 contiguous f32 activations per block, so
16// packing is the kernel's requirement rather than a memory optimization.
17//
18// Weights: projections read as RAW Q8_0 (int8 + per-block f16 scale, pre-decoded once); norms and
19// the adaLN vector read as f32.
20// license_tier: ORIGINAL
21
22import "nx_syscalls.nx"
23import "nx_le.nx"
24import "nx_f32.nx"
25import "nx_f32_div.nx"
26import "nx_f32_cvt.nx"
27import "nx_f16.nx"
28import "nx_f32_exp.nx"
29import "nx_f32_activations.nx"
30import "nx_strconv.nx"
31import "nx_genfix.nx"
32import "nx_genver.nx"
33import "nx_genweights.nx"
34
35import "nx_genblock.nx"
36const K_MAGIC_1000000: i64 = 1000000
37
38func main(argc: i64, argv: *i64) -> i64 {
39 if argc < 2 {
40 br_puts("usage: nx_gen_blockrun <model_id> [n_workers] [reference_name]\n" as *u8)
41 return 2
42 }
43 let M: *u8 = argv[1] as *u8
44 var NW: i64 = 16
45 if argc >= 3 {
46 let ep: *i64 = sys_mmap(32) as *i64
47 ep[0] = 0
48 NW = nx_strconv_parse_i64(argv[2] as *u8, ep)
49 if ep[0] != 0 { NW = 16 }
50 if NW < 1 { NW = 1 }
51 }
52
53 // argv[4] = GGUF path. Present -> weights come from the model file; absent -> from the
54 // dumped fixtures. Both paths kept so the swap can be A/B'd instead of trusted.
55 var gw: *i64 = 0 as *i64
56 if argc >= 5 {
57 gw = nx_gw_open(argv[4] as *u8)
58 if (gw as i64) == 0 { br_puts("gguf open/parse failed
59" as *u8); return 20 }
60 nx_genver_emit("gguf_tensors" as *u8, nx_gw_ntensors(gw))
61 }
62
63 let ne: *i64 = sys_mmap(64) as *i64
64 let n_in: *u8 = "layers.0.blk_in" as *u8
65 let c_in: i64 = nx_genfix_dims(M, n_in, br_strlen(n_in), ne)
66 if c_in < 0 { br_puts("missing layers.0.blk_in\n" as *u8); return 30 }
67 let D: i64 = ne[0]
68 let T: i64 = ne[1]
69 let QKV: i64 = (BR_N_HEADS * 3) * BR_HEAD_DIM
70 let ne2: *i64 = sys_mmap(64) as *i64
71 let n_w1: *u8 = "model.diffusion_model.layers.0.feed_forward.w1.weight" as *u8
72 if nx_genfix_dims(M, n_w1, br_strlen(n_w1), ne2) < 0 { br_puts("missing w1\n" as *u8); return 31 }
73 let FD: i64 = ne2[1]
74
75 nx_genver_emit("dim" as *u8, D)
76 nx_genver_emit("tokens" as *u8, T)
77 nx_genver_emit("qkv_width" as *u8, QKV)
78 nx_genver_emit("ffn_dim" as *u8, FD)
79 nx_genver_emit("workers" as *u8, NW)
80
81 // ---- inputs -----------------------------------------------------------------------
82 let blk_in: *u8 = nx_genfix_load(M, n_in, br_strlen(n_in), c_in)
83 if (blk_in as i64) == 0 { br_puts("load blk_in failed\n" as *u8); return 40 }
84 let n_ad: *u8 = "layers.0.adaln_m" as *u8
85 let adaln: *u8 = nx_genfix_load(M, n_ad, br_strlen(n_ad), BR_CHUNKS * D)
86 if (adaln as i64) == 0 { br_puts("load adaln failed\n" as *u8); return 41 }
87 // SELF-CHECK: recompute layer 0's adaLN from the GGUF and compare against the tapped
88 // fixture. Chaining N layers requires computing this per layer, and a wrong adaLN produces
89 // plausible-but-wrong activations rather than an error -- so prove it on the one layer where
90 // an oracle copy exists before relying on it for the rest.
91 if (gw as i64) != 0 {
92 let n_te: *u8 = "t_emb" as *u8
93 let t_emb: *u8 = nx_genfix_load(M, n_te, br_strlen(n_te), 256)
94 if (t_emb as i64) != 0 {
95 let calc: *u8 = sys_mmap(BR_CHUNKS * D * 4 + 64)
96 let rc_a: i64 = br_adaln(gw, 0, t_emb, calc, 256, BR_CHUNKS * D)
97 nx_genver_emit("adaln_from_gguf_rc" as *u8, rc_a)
98 if rc_a == 0 {
99 let tolq: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(1000))
100 var bad: i64 = 0
101 var ai: i64 = 0
102 while ai < BR_CHUNKS * D {
103 if nx_genver_exceeds(nx_le_read_u32(calc, ai * 4),
104 nx_le_read_u32(adaln, ai * 4), tolq) == 1 { bad = bad + 1 }
105 ai = ai + 1
106 }
107 nx_genver_emit("adaln_gguf_vs_fixture_fail_1e3" as *u8, bad)
108 }
109 }
110 }
111
112 let n_pe: *u8 = "pe" as *u8
113 let pe: *u8 = nx_genfix_load(M, n_pe, br_strlen(n_pe), 2 * 2 * (BR_HEAD_DIM / 2) * T)
114 if (pe as i64) == 0 { br_puts("load pe failed\n" as *u8); return 42 }
115
116 var wn1: *u8 = 0 as *u8
117 if (gw as i64) != 0 { wn1 = br_gw_f32(gw, "model.diffusion_model.layers.0.attention_norm1.weight" as *u8, D) }
118 if (gw as i64) == 0 { wn1 = nx_genfix_load(M, "model.diffusion_model.layers.0.attention_norm1.weight" as *u8, br_strlen("model.diffusion_model.layers.0.attention_norm1.weight" as *u8), D) }
119 var wn2: *u8 = 0 as *u8
120 if (gw as i64) != 0 { wn2 = br_gw_f32(gw, "model.diffusion_model.layers.0.attention_norm2.weight" as *u8, D) }
121 if (gw as i64) == 0 { wn2 = nx_genfix_load(M, "model.diffusion_model.layers.0.attention_norm2.weight" as *u8, br_strlen("model.diffusion_model.layers.0.attention_norm2.weight" as *u8), D) }
122 var wf1: *u8 = 0 as *u8
123 if (gw as i64) != 0 { wf1 = br_gw_f32(gw, "model.diffusion_model.layers.0.ffn_norm1.weight" as *u8, D) }
124 if (gw as i64) == 0 { wf1 = nx_genfix_load(M, "model.diffusion_model.layers.0.ffn_norm1.weight" as *u8, br_strlen("model.diffusion_model.layers.0.ffn_norm1.weight" as *u8), D) }
125 var wf2: *u8 = 0 as *u8
126 if (gw as i64) != 0 { wf2 = br_gw_f32(gw, "model.diffusion_model.layers.0.ffn_norm2.weight" as *u8, D) }
127 if (gw as i64) == 0 { wf2 = nx_genfix_load(M, "model.diffusion_model.layers.0.ffn_norm2.weight" as *u8, br_strlen("model.diffusion_model.layers.0.ffn_norm2.weight" as *u8), D) }
128 var wqn: *u8 = 0 as *u8
129 if (gw as i64) != 0 { wqn = br_gw_f32(gw, "model.diffusion_model.layers.0.attention.q_norm.weight" as *u8, BR_HEAD_DIM) }
130 if (gw as i64) == 0 { wqn = nx_genfix_load(M, "model.diffusion_model.layers.0.attention.q_norm.weight" as *u8, br_strlen("model.diffusion_model.layers.0.attention.q_norm.weight" as *u8), BR_HEAD_DIM) }
131 var wkn: *u8 = 0 as *u8
132 if (gw as i64) != 0 { wkn = br_gw_f32(gw, "model.diffusion_model.layers.0.attention.k_norm.weight" as *u8, BR_HEAD_DIM) }
133 if (gw as i64) == 0 { wkn = nx_genfix_load(M, "model.diffusion_model.layers.0.attention.k_norm.weight" as *u8, br_strlen("model.diffusion_model.layers.0.attention.k_norm.weight" as *u8), BR_HEAD_DIM) }
134 if (wn1 as i64) == 0 { br_puts("load norm1 failed\n" as *u8); return 43 }
135 if (wn2 as i64) == 0 { br_puts("load norm2 failed\n" as *u8); return 44 }
136 if (wf1 as i64) == 0 { br_puts("load ffn_norm1 failed\n" as *u8); return 45 }
137 if (wf2 as i64) == 0 { br_puts("load ffn_norm2 failed\n" as *u8); return 46 }
138 if (wqn as i64) == 0 { br_puts("load q_norm failed\n" as *u8); return 47 }
139 if (wkn as i64) == 0 { br_puts("load k_norm failed\n" as *u8); return 48 }
140
141 let s_qkv: *i64 = sys_mmap((D / 32) * QKV * 8) as *i64
142 var W_qkv: *u8 = 0 as *u8
143 if (gw as i64) != 0 { W_qkv = br_gw_q8(gw, "model.diffusion_model.layers.0.attention.qkv.weight" as *u8, D, QKV, s_qkv) }
144 if (gw as i64) == 0 { W_qkv = br_load_q8(M, "model.diffusion_model.layers.0.attention.qkv.weight" as *u8, D, QKV, s_qkv) }
145 if (W_qkv as i64) == 0 { br_puts("load qkv weight failed\n" as *u8); return 50 }
146 let s_out: *i64 = sys_mmap((D / 32) * D * 8) as *i64
147 var W_out: *u8 = 0 as *u8
148 if (gw as i64) != 0 { W_out = br_gw_q8(gw, "model.diffusion_model.layers.0.attention.out.weight" as *u8, D, D, s_out) }
149 if (gw as i64) == 0 { W_out = br_load_q8(M, "model.diffusion_model.layers.0.attention.out.weight" as *u8, D, D, s_out) }
150 if (W_out as i64) == 0 { br_puts("load out weight failed\n" as *u8); return 51 }
151 let s_w1: *i64 = sys_mmap((D / 32) * FD * 8) as *i64
152 var W_w1: *u8 = 0 as *u8
153 if (gw as i64) != 0 { W_w1 = br_gw_q8(gw, n_w1, D, FD, s_w1) }
154 if (gw as i64) == 0 { W_w1 = br_load_q8(M, n_w1, D, FD, s_w1) }
155 if (W_w1 as i64) == 0 { br_puts("load w1 failed\n" as *u8); return 52 }
156 let s_w3: *i64 = sys_mmap((D / 32) * FD * 8) as *i64
157 var W_w3: *u8 = 0 as *u8
158 if (gw as i64) != 0 { W_w3 = br_gw_q8(gw, "model.diffusion_model.layers.0.feed_forward.w3.weight" as *u8, D, FD, s_w3) }
159 if (gw as i64) == 0 { W_w3 = br_load_q8(M, "model.diffusion_model.layers.0.feed_forward.w3.weight" as *u8, D, FD, s_w3) }
160 if (W_w3 as i64) == 0 { br_puts("load w3 failed\n" as *u8); return 53 }
161 let s_w2: *i64 = sys_mmap((FD / 32) * D * 8) as *i64
162 var W_w2: *u8 = 0 as *u8
163 if (gw as i64) != 0 { W_w2 = br_gw_q8(gw, "model.diffusion_model.layers.0.feed_forward.w2.weight" as *u8, FD, D, s_w2) }
164 if (gw as i64) == 0 { W_w2 = br_load_q8(M, "model.diffusion_model.layers.0.feed_forward.w2.weight" as *u8, FD, D, s_w2) }
165 if (W_w2 as i64) == 0 { br_puts("load w2 failed\n" as *u8); return 54 }
166
167 // ---- scratch (SHARED: forked matmul workers write into these) ----------------------
168 let h_attn: *u8 = sys_mmap_shared(T * D * 4)
169 let qkv: *u8 = sys_mmap_shared(T * QKV * 4)
170 let qrp: *u8 = sys_mmap_shared(T * BR_N_HEADS * BR_HEAD_DIM * 4)
171 let krp: *u8 = sys_mmap_shared(T * BR_N_HEADS * BR_HEAD_DIM * 4)
172 let aop: *u8 = sys_mmap_shared(T * D * 4)
173 let atto: *u8 = sys_mmap_shared(T * D * 4)
174 let an2: *u8 = sys_mmap_shared(T * D * 4)
175 let mid: *u8 = sys_mmap_shared(T * D * 4)
176 let h_ffn: *u8 = sys_mmap_shared(T * D * 4)
177 let fw1: *u8 = sys_mmap_shared(T * FD * 4)
178 let fw3: *u8 = sys_mmap_shared(T * FD * 4)
179 let fact: *u8 = sys_mmap_shared(T * FD * 4)
180 let fdn: *u8 = sys_mmap_shared(T * D * 4)
181 let fn2: *u8 = sys_mmap_shared(T * D * 4)
182 let outb: *u8 = sys_mmap_shared(T * D * 4)
183
184 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(K_MAGIC_1000000))
185 let one: i64 = nx_i32_to_f32(1)
186
187 // adaLN chunks: [scale_msa | gate_msa | scale_mlp | gate_mlp]
188 let mod_msa: *i64 = sys_mmap(D * 8) as *i64
189 let tg_msa: *i64 = sys_mmap(D * 8) as *i64
190 let mod_mlp: *i64 = sys_mmap(D * 8) as *i64
191 let tg_mlp: *i64 = sys_mmap(D * 8) as *i64
192 var i0: i64 = 0
193 while i0 < D {
194 mod_msa[i0] = __f32_add(one, nx_le_read_u32(adaln, (0 * D + i0) * 4))
195 tg_msa[i0] = nx_f32_tanh(nx_le_read_u32(adaln, (1 * D + i0) * 4))
196 mod_mlp[i0] = __f32_add(one, nx_le_read_u32(adaln, (2 * D + i0) * 4))
197 tg_mlp[i0] = nx_f32_tanh(nx_le_read_u32(adaln, (3 * D + i0) * 4))
198 i0 = i0 + 1
199 }
200
201 let t0: i64 = sys_now_us()
202
203 // 1. pre-norm + modulate
204 br_rmsnorm(blk_in, wn1, mod_msa, h_attn, T, D, eps)
205 // 2. qkv projection
206 br_matmul(W_qkv, s_qkv, h_attn, qkv, T, D, QKV, NW)
207 let t_qkv: i64 = sys_now_us()
208
209 // 3+4. qk-norm then RoPE, writing straight into the post-RoPE layout [head_dim, L, n_heads]
210 let half: i64 = BR_HEAD_DIM / 2
211 let hdinv: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(BR_HEAD_DIM))
212 var l: i64 = 0
213 while l < T {
214 var h: i64 = 0
215 while h < BR_N_HEADS {
216 var side: i64 = 0
217 while side < 2 {
218 var wnorm: *u8 = wqn
219 var dst: *u8 = qrp
220 var hbase: i64 = 0
221 if side == 1 { wnorm = wkn; dst = krp; hbase = BR_N_HEADS }
222 let src: i64 = l * QKV + (hbase + h) * BR_HEAD_DIM
223 var ss: i64 = 0
224 var d: i64 = 0
225 while d < BR_HEAD_DIM {
226 let v: i64 = nx_le_read_u32(qkv, (src + d) * 4)
227 ss = __f32_add(ss, __f32_mul(v, v))
228 d = d + 1
229 }
230 let inv: i64 = nx_f32_div(nx_i32_to_f32(1),
231 nx_f32_sqrt(__f32_add(__f32_mul(ss, hdinv), eps)))
232 let orow: i64 = (h * T + l) * BR_HEAD_DIM
233 var j: i64 = 0
234 while j < half {
235 let a0: i64 = __f32_mul(__f32_mul(nx_le_read_u32(qkv, (src + 2 * j) * 4), inv),
236 nx_le_read_u32(wnorm, (2 * j) * 4))
237 let a1: i64 = __f32_mul(__f32_mul(nx_le_read_u32(qkv, (src + 2 * j + 1) * 4), inv),
238 nx_le_read_u32(wnorm, (2 * j + 1) * 4))
239 var r: i64 = 0
240 while r < 2 {
241 let pb: i64 = ((l * half + j) * 2 + r) * 2
242 let v: i64 = __f32_add(__f32_mul(a0, nx_le_read_u32(pe, pb * 4)),
243 __f32_mul(a1, nx_le_read_u32(pe, (pb + 1) * 4)))
244 nx_le_write_u32(dst, (orow + 2 * j + r) * 4, v)
245 r = r + 1
246 }
247 j = j + 1
248 }
249 side = side + 1
250 }
251 h = h + 1
252 }
253 l = l + 1
254 }
255
256 let t_rope: i64 = sys_now_us()
257 // 5. SDPA -- forked over heads
258 let scale: i64 = nx_f32_div(nx_i32_to_f32(1), nx_f32_sqrt(nx_i32_to_f32(BR_HEAD_DIM)))
259 var sw: i64 = NW
260 if sw > BR_N_HEADS { sw = BR_N_HEADS }
261 if sw <= 1 {
262 br_sdpa_band(qrp, krp, qkv, aop, T, D, QKV, BR_HEAD_DIM, BR_N_HEADS, scale, 0, BR_N_HEADS)
263 }
264 if sw > 1 {
265 let spids: *i64 = sys_mmap(sw * 8 + 64) as *i64
266 var sk: i64 = 0
267 while sk < sw {
268 let a0: i64 = sk * BR_N_HEADS / sw
269 let a1: i64 = (sk + 1) * BR_N_HEADS / sw
270 let pid2: i64 = sys_fork()
271 if pid2 == 0 {
272 br_sdpa_band(qrp, krp, qkv, aop, T, D, QKV, BR_HEAD_DIM, BR_N_HEADS, scale, a0, a1)
273 sys_exit(0)
274 }
275 spids[sk] = pid2
276 sk = sk + 1
277 }
278 let sst: *i64 = sys_mmap(64) as *i64
279 sk = 0
280 while sk < sw { sys_wait4(spids[sk], sst, 0); sk = sk + 1 }
281 }
282
283 let t_sdpa: i64 = sys_now_us()
284 // 6..8 attention tail
285 br_matmul(W_out, s_out, aop, atto, T, D, D, NW)
286 br_rmsnorm(atto, wn2, 0 as *i64, an2, T, D, eps)
287 br_gate_resid(an2, tg_msa, blk_in, mid, T, D)
288
289 let t_tail: i64 = sys_now_us()
290 // 9..14 feed-forward half
291 br_rmsnorm(mid, wf1, mod_mlp, h_ffn, T, D, eps)
292 br_matmul(W_w1, s_w1, h_ffn, fw1, T, D, FD, NW)
293 br_matmul(W_w3, s_w3, h_ffn, fw3, T, D, FD, NW)
294 br_fork_range(0, fw1, fw3, fact, T * FD, NW)
295 br_matmul(W_w2, s_w2, fact, fdn, T, FD, D, NW)
296 br_rmsnorm(fdn, wf2, 0 as *i64, fn2, T, D, eps)
297 br_gate_resid(fn2, tg_mlp, mid, outb, T, D)
298
299 let t1: i64 = sys_now_us()
300 nx_genver_emit("block_us" as *u8, t1 - t0)
301 // Per-stage split: guessing which stage dominates is how the scalar SDPA hid behind tuned
302 // matmuls for a whole round. Measure the phases, then optimize the biggest one.
303 nx_genver_emit("us_prenorm_qkv" as *u8, t_qkv - t0)
304 nx_genver_emit("us_qknorm_rope" as *u8, t_rope - t_qkv)
305 nx_genver_emit("us_sdpa" as *u8, t_sdpa - t_rope)
306 nx_genver_emit("us_attn_tail" as *u8, t_tail - t_sdpa)
307 nx_genver_emit("us_ffn" as *u8, t1 - t_tail)
308
309 // ---- grade the assembled output against the oracle's next-block input --------------
310 // The reference is an ARGUMENT: grading against the oracle's next-block input measures our
311 // error PLUS the oracle's own quantization error, and cannot separate them. Grading against
312 // an f64 reference of the same block answers "did our chain stay on the rails" on its own.
313 var n_ref: *u8 = "layers.1.blk_in" as *u8
314 if argc >= 4 { n_ref = argv[3] as *u8 }
315 let c_ref: i64 = nx_genfix_dims(M, n_ref, br_strlen(n_ref), ne)
316 if c_ref < 0 { br_puts("missing layers.1.blk_in\n" as *u8); return 60 }
317 let ref: *u8 = nx_genfix_load(M, n_ref, br_strlen(n_ref), c_ref)
318 if (ref as i64) == 0 { br_puts("load reference failed\n" as *u8); return 61 }
319
320 let tol: *i64 = sys_mmap(64) as *i64
321 nx_genver_tols(tol)
322 let c: *i64 = sys_mmap(128) as *i64
323 nx_genver_init(c, 1)
324 var f: i64 = 0
325 while f < T * D {
326 nx_genver_tally(c, tol, nx_le_read_u32(outb, f * 4), nx_le_read_u32(ref, f * 4), f)
327 f = f + 1
328 }
329 return nx_genver_report(c)
330}