code wiki / (root) / nx_gen_blockrun.nx

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}