code wiki / (root) / nx_gen_ditfull.nx

nx_gen_ditfull.nx source

↩ module page · 379 lines · 16535 B

1// nx_gen_ditfull.nx -- the COMPLETE sovereign Z-Image DiT forward, weights from the GGUF. 2// 3// txt_embed --[context_refiner x2, UNMODULATED]--> txt 4// img_embed --[noise_refiner x2, modulated ]--> img 5// concat(txt, img) --[layers x30, modulated]--> h --[final layer]--> [64, L] 6// 7// Everything after the text encoder and the latent embedding. No oracle tensor in the loop: 8// weights are read from the model file, adaLN is computed per block, and every activation is the 9// previous stage's own output. 10// 11// Usage: nx_gen_ditfull <model_id> <gguf_path> [workers] [reference] 12// 13// THE FINAL LAYER IS NOT A BLOCK and does not reuse br_block: 14// scale = W_ada1 . silu(t_emb) + b <- silu HERE, unlike the blocks' adaLN 15// x = LayerNorm(x) <- LayerNorm (mean-centred), NOT RMSNorm, no affine 16// x = x * (1 + scale) 17// x = W_lin . x + b 18// Both differences were read out of FinalLayer::forward, not assumed. Reusing the block's adaLN 19// path here would apply no silu and subtract no mean, and would still produce a finite tensor. 20// 21// ⚠TOKEN PADDING IS NOT IMPLEMENTED: this shape needs none (512 text and 256 image tokens are both 22// multiples of SEQ_MULTI_OF=32, so n_pad=0). A shape that needs padding must append the learned 23// cap_pad_token / x_pad_token first; this organ REFUSES rather than silently skipping it. 24// license_tier: ORIGINAL 25 26import "nx_syscalls.nx" 27import "nx_le.nx" 28import "nx_f32.nx" 29import "nx_f32_div.nx" 30import "nx_f32_cvt.nx" 31import "nx_f16.nx" 32import "nx_f32_exp.nx" 33import "nx_f32_activations.nx" 34import "nx_strconv.nx" 35import "nx_genfix.nx" 36import "nx_genver.nx" 37import "nx_genweights.nx" 38import "nx_genblock.nx" 39 40const DF_SEQ_MULTI: i64 = 32 41const DF_EMB: i64 = 256 42 43func df_puts(s: *u8) -> i64 { 44 var n: i64 = 0 45 while s[n] != (0 as u8) { n = n + 1 } 46 return sys_write(1, s, n) 47} 48 49// ---- final layer ----------------------------------------------------------------------- 50func df_final(gw: *i64, x: *u8, t_emb: *u8, out: *u8, T: i64, D: i64, OUTD: i64, NW: i64) -> i64 { 51 // scale = W_ada1 . silu(t_emb) + b 52 let se: *u8 = sys_mmap(DF_EMB * 4 + 64) 53 var i: i64 = 0 54 while i < DF_EMB { 55 nx_le_write_u32(se, i * 4, nx_f32_silu(nx_le_read_u32(t_emb, i * 4))) 56 i = i + 1 57 } 58 let wn: *u8 = "model.diffusion_model.final_layer.adaLN_modulation.1.weight" as *u8 59 // F16 in the file -> materialize once as packed f32 (see br_matmul_f32's note) 60 let w: *u8 = br_gw_f32(gw, wn, DF_EMB * D) 61 if (w as i64) == 0 { return 0 - 1 } 62 let bn: *u8 = "model.diffusion_model.final_layer.adaLN_modulation.1.bias" as *u8 63 let bias: *u8 = br_gw_f32(gw, bn, D) 64 if (bias as i64) == 0 { return 0 - 2 } 65 66 let scale: *u8 = sys_mmap(D * 4 + 64) 67 br_matmul_f32(w, se, scale, 1, DF_EMB, D, 1) 68 var o: i64 = 0 69 while o < D { 70 nx_le_write_u32(scale, o * 4, 71 __f32_add(nx_le_read_u32(scale, o * 4), nx_le_read_u32(bias, o * 4))) 72 o = o + 1 73 } 74 75 // LayerNorm (no affine) then modulate, per token 76 let one: i64 = nx_i32_to_f32(1) 77 let dinv: i64 = nx_f32_div(one, nx_i32_to_f32(D)) 78 let eps: i64 = nx_f32_div(one, nx_i32_to_f32(1000000)) 79 let tmp: *u8 = sys_mmap(T * D * 4 + 64) 80 var t: i64 = 0 81 while t < T { 82 let base: i64 = t * D 83 var mean: i64 = 0 84 var k: i64 = 0 85 while k < D { mean = __f32_add(mean, nx_le_read_u32(x, (base + k) * 4)); k = k + 1 } 86 mean = __f32_mul(mean, dinv) 87 let nmean: i64 = mean ^ 0x80000000 88 var vr: i64 = 0 89 k = 0 90 while k < D { 91 let d: i64 = __f32_add(nx_le_read_u32(x, (base + k) * 4), nmean) 92 vr = __f32_add(vr, __f32_mul(d, d)) 93 k = k + 1 94 } 95 let inv: i64 = nx_f32_div(one, nx_f32_sqrt(__f32_add(__f32_mul(vr, dinv), eps))) 96 k = 0 97 while k < D { 98 let d: i64 = __f32_add(nx_le_read_u32(x, (base + k) * 4), nmean) 99 let v: i64 = __f32_mul(__f32_mul(d, inv), 100 __f32_add(one, nx_le_read_u32(scale, k * 4))) 101 nx_le_write_u32(tmp, (base + k) * 4, v) 102 k = k + 1 103 } 104 t = t + 1 105 } 106 107 // linear (with bias) -> [OUTD, T] 108 let lw: *u8 = "model.diffusion_model.final_layer.linear.weight" as *u8 109 let W: *u8 = br_gw_f32(gw, lw, D * OUTD) 110 if (W as i64) == 0 { return 0 - 3 } 111 let lb: *u8 = "model.diffusion_model.final_layer.linear.bias" as *u8 112 let lbias: *u8 = br_gw_f32(gw, lb, OUTD) 113 if (lbias as i64) == 0 { return 0 - 4 } 114 br_matmul_f32(W, tmp, out, T, D, OUTD, NW) 115 var tt: i64 = 0 116 while tt < T { 117 var oo: i64 = 0 118 while oo < OUTD { 119 let f: i64 = tt * OUTD + oo 120 nx_le_write_u32(out, f * 4, 121 __f32_add(nx_le_read_u32(out, f * 4), nx_le_read_u32(lbias, oo * 4))) 122 oo = oo + 1 123 } 124 tt = tt + 1 125 } 126 return 0 127} 128 129func main(argc: i64, argv: *i64) -> i64 { 130 if argc < 3 { 131 df_puts("usage: nx_gen_ditfull <model_id> <gguf_path> [workers] [reference]\n" as *u8) 132 return 2 133 } 134 let M: *u8 = argv[1] as *u8 135 let GP: *u8 = argv[2] as *u8 136 var NW: i64 = 16 137 if argc >= 4 { 138 let ep: *i64 = sys_mmap(32) as *i64 139 ep[0] = 0 140 NW = nx_strconv_parse_i64(argv[3] as *u8, ep) 141 if ep[0] != 0 { NW = 16 } 142 if NW < 1 { NW = 1 } 143 } 144 145 let gw: *i64 = nx_gw_open(GP) 146 if (gw as i64) == 0 { df_puts("gguf open failed\n" as *u8); return 20 } 147 // Bind the geometry from THIS checkpoint's own tensor shapes before any block runs. 148 // A hardcoded head count is a model welded into the binary; probing makes it swappable. 149 if br_arch_bind(gw) != 0 { df_puts("architecture not derivable from this checkpoint -- refusing 150" as *u8); return 21 } 151 152 let ne: *i64 = sys_mmap(64) as *i64 153 let n_tx: *u8 = "txt_embed" as *u8 154 let c_tx: i64 = nx_genfix_dims(M, n_tx, br_strlen(n_tx), ne) 155 if c_tx < 0 { df_puts("missing txt_embed\n" as *u8); return 30 } 156 let D: i64 = ne[0] 157 let TT: i64 = ne[1] 158 let n_im: *u8 = "img_embed" as *u8 159 let c_im: i64 = nx_genfix_dims(M, n_im, br_strlen(n_im), ne) 160 if c_im < 0 { df_puts("missing img_embed\n" as *u8); return 31 } 161 let TI: i64 = ne[1] 162 let T: i64 = TT + TI 163 let QKV: i64 = (BR_N_HEADS * 3) * BR_HEAD_DIM 164 165 // Refuse rather than silently skip the learned pad tokens. 166 if TT - (TT / DF_SEQ_MULTI) * DF_SEQ_MULTI != 0 { df_puts("txt tokens need padding: unimplemented\n" as *u8); return 32 } 167 if TI - (TI / DF_SEQ_MULTI) * DF_SEQ_MULTI != 0 { df_puts("img tokens need padding: unimplemented\n" as *u8); return 33 } 168 169 let nm0: *u8 = sys_mmap(256) 170 let wi: i64 = nx_gw_find(gw, br_name_p(nm0, "layers" as *u8, 0, "feed_forward.w1.weight" as *u8), br_strlen(nm0)) 171 if wi < 0 { df_puts("missing layers.0 w1\n" as *u8); return 34 } 172 let FD: i64 = nx_gw_dim1(gw, wi) 173 174 nx_genver_emit("dim" as *u8, D) 175 nx_genver_emit("txt_tokens" as *u8, TT) 176 nx_genver_emit("img_tokens" as *u8, TI) 177 nx_genver_emit("tokens" as *u8, T) 178 nx_genver_emit("ffn_dim" as *u8, FD) 179 nx_genver_emit("workers" as *u8, NW) 180 181 let txt0: *u8 = nx_genfix_load(M, n_tx, br_strlen(n_tx), c_tx) 182 if (txt0 as i64) == 0 { df_puts("load txt_embed failed\n" as *u8); return 40 } 183 let img0: *u8 = nx_genfix_load(M, n_im, br_strlen(n_im), c_im) 184 if (img0 as i64) == 0 { df_puts("load img_embed failed\n" as *u8); return 41 } 185 let n_te: *u8 = "t_emb" as *u8 186 let t_emb: *u8 = nx_genfix_load(M, n_te, br_strlen(n_te), DF_EMB) 187 if (t_emb as i64) == 0 { df_puts("load t_emb failed\n" as *u8); return 42 } 188 let n_pe: *u8 = "pe" as *u8 189 let pe: *u8 = nx_genfix_load(M, n_pe, br_strlen(n_pe), 2 * 2 * (BR_HEAD_DIM / 2) * T) 190 if (pe as i64) == 0 { df_puts("load pe failed\n" as *u8); return 43 } 191 // pe rows are per token: txt uses the first TT, img the next TI. 192 let pe_img: *u8 = ((pe as i64) + TT * (BR_HEAD_DIM / 2) * 4 * 4) as *u8 193 194 // scratch sized for the WIDEST stage (the concatenated 768-token stack) 195 let scr: *i64 = sys_mmap(BR_SCR_SLOTS * 8 + 64) as *i64 196 scr[0] = sys_mmap_shared(T * D * 4) as i64 197 scr[1] = sys_mmap_shared(T * QKV * 4) as i64 198 scr[2] = sys_mmap_shared(T * BR_N_HEADS * BR_HEAD_DIM * 4) as i64 199 scr[3] = sys_mmap_shared(T * BR_N_HEADS * BR_HEAD_DIM * 4) as i64 200 scr[4] = sys_mmap_shared(T * D * 4) as i64 201 scr[5] = sys_mmap_shared(T * D * 4) as i64 202 scr[6] = sys_mmap_shared(T * D * 4) as i64 203 scr[7] = sys_mmap_shared(T * D * 4) as i64 204 scr[8] = sys_mmap_shared(T * D * 4) as i64 205 scr[9] = sys_mmap_shared(T * FD * 4) as i64 206 scr[10] = sys_mmap_shared(T * FD * 4) as i64 207 scr[11] = sys_mmap_shared(T * FD * 4) as i64 208 scr[12] = sys_mmap_shared(T * D * 4) as i64 209 scr[13] = sys_mmap_shared(T * D * 4) as i64 210 scr[14] = sys_mmap((D / 32) * QKV * 8) as i64 211 scr[15] = sys_mmap((D / 32) * D * 8) as i64 212 scr[16] = sys_mmap((D / 32) * FD * 8) as i64 213 scr[17] = sys_mmap((D / 32) * FD * 8) as i64 214 scr[18] = sys_mmap((FD / 32) * D * 8) as i64 215 scr[19] = sys_mmap(BR_CHUNKS * D * 4 + 64) as i64 216 scr[20] = sys_mmap(D * 8) as i64 217 scr[21] = sys_mmap(D * 8) as i64 218 scr[22] = sys_mmap(D * 8) as i64 219 scr[23] = sys_mmap(D * 8) as i64 220 221 let bufA: *u8 = sys_mmap_shared(T * D * 4) 222 let bufB: *u8 = sys_mmap_shared(T * D * 4) 223 let cat: *u8 = sys_mmap_shared(T * D * 4) 224 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(100000)) // norm_eps = 1e-5 225 226 let t0: i64 = sys_now_us() 227 228 // ---- context_refiner x2 on the text stream (UNMODULATED) -------------------------- 229 var cur: *u8 = txt0 230 var nxt: *u8 = bufA 231 var i: i64 = 0 232 while i < 2 { 233 let rc: i64 = br_block(gw, "context_refiner" as *u8, i, 0, cur, nxt, scr, 234 TT, D, QKV, FD, NW, pe, t_emb, eps) 235 if rc != 0 { nx_genver_emit("context_refiner_failed" as *u8, i); nx_genver_emit("rc" as *u8, rc); return 50 } 236 cur = nxt 237 if (nxt as i64) == (bufA as i64) { nxt = bufB } else { nxt = bufA } 238 i = i + 1 239 } 240 let txt_done: *u8 = cur 241 let t_ctx: i64 = sys_now_us() 242 243 // ---- noise_refiner x2 on the image stream (MODULATED) ----------------------------- 244 let bufC: *u8 = sys_mmap_shared(T * D * 4) 245 let bufD: *u8 = sys_mmap_shared(T * D * 4) 246 cur = img0 247 nxt = bufC 248 i = 0 249 while i < 2 { 250 let rc: i64 = br_block(gw, "noise_refiner" as *u8, i, 1, cur, nxt, scr, 251 TI, D, QKV, FD, NW, pe_img, t_emb, eps) 252 if rc != 0 { nx_genver_emit("noise_refiner_failed" as *u8, i); nx_genver_emit("rc" as *u8, rc); return 51 } 253 cur = nxt 254 if (nxt as i64) == (bufC as i64) { nxt = bufD } else { nxt = bufC } 255 i = i + 1 256 } 257 let img_done: *u8 = cur 258 let t_noise: i64 = sys_now_us() 259 260 // ---- concat(txt, img) -------------------------------------------------------------- 261 var e: i64 = 0 262 while e < TT * D { nx_le_write_u32(cat, e * 4, nx_le_read_u32(txt_done, e * 4)); e = e + 1 } 263 e = 0 264 while e < TI * D { nx_le_write_u32(cat, (TT * D + e) * 4, nx_le_read_u32(img_done, e * 4)); e = e + 1 } 265 266 // ---- layers x30 -------------------------------------------------------------------- 267 cur = cat 268 nxt = bufA 269 var lay: i64 = 0 270 while lay < 30 { 271 let rc: i64 = br_block(gw, "layers" as *u8, lay, 1, cur, nxt, scr, 272 T, D, QKV, FD, NW, pe, t_emb, eps) 273 if rc != 0 { nx_genver_emit("layer_failed" as *u8, lay); nx_genver_emit("rc" as *u8, rc); return 52 } 274 cur = nxt 275 if (nxt as i64) == (bufA as i64) { nxt = bufB } else { nxt = bufA } 276 lay = lay + 1 277 } 278 let t_layers: i64 = sys_now_us() 279 280 // ---- final layer ------------------------------------------------------------------- 281 let OUTD: i64 = 64 282 let fout: *u8 = sys_mmap_shared(T * OUTD * 4 + 64) 283 let rcf: i64 = df_final(gw, cur, t_emb, fout, T, D, OUTD, NW) 284 if rcf != 0 { nx_genver_emit("final_layer_failed" as *u8, rcf); return 53 } 285 let t1: i64 = sys_now_us() 286 287 nx_genver_emit("us_context_refiner" as *u8, t_ctx - t0) 288 nx_genver_emit("us_noise_refiner" as *u8, t_noise - t_ctx) 289 nx_genver_emit("us_layers" as *u8, t_layers - t_noise) 290 nx_genver_emit("us_final" as *u8, t1 - t_layers) 291 nx_genver_emit("dit_us" as *u8, t1 - t0) 292 293 // ---- bridge to the decoder: unpatchify -> sampler step -> tanh pre-scale -------------- 294 // 295 // latent[c][2hh+i][2ww+j] = -token[TT + hh*gw + ww][(i*p + j)*C + c] (unpatchify + NEGATE) 296 // z = x0 - 1.0 * latent (one Euler step, sigma0 = 1) 297 // tae_h = tanh(z/3) * 3 (the TAESD decoder's entry activation) 298 // 299 // The negation and sigma0=1 were both MEASURED, not read off a formula: unpatchify matched to 300 // exactly 0 only with the sign flipped, and the implied x0 is unit gaussian (mean 0.003, 301 // std 0.997) only for v = +dit_out. A wrong sign here still decodes to a picture -- just the 302 // wrong one -- so it had to be pinned by evidence. 303 let PATCH: i64 = 2 304 let LC: i64 = 16 305 let GH: i64 = 16 306 let GWD: i64 = 16 307 let LH: i64 = GH * PATCH 308 let LW: i64 = GWD * PATCH 309 let lat: *u8 = sys_mmap(LC * LH * LW * 4 + 64) 310 var hh: i64 = 0 311 while hh < GH { 312 var ww: i64 = 0 313 while ww < GWD { 314 let tok: i64 = (TT + hh * GWD + ww) * OUTD 315 var ii: i64 = 0 316 while ii < PATCH { 317 var jj: i64 = 0 318 while jj < PATCH { 319 let pp: i64 = (ii * PATCH + jj) * LC 320 var cc2: i64 = 0 321 while cc2 < LC { 322 let v: i64 = nx_le_read_u32(fout, (tok + pp + cc2) * 4) ^ 0x80000000 323 let y: i64 = hh * PATCH + ii 324 let xx: i64 = ww * PATCH + jj 325 nx_le_write_u32(lat, (cc2 * LH * LW + y * LW + xx) * 4, v) 326 cc2 = cc2 + 1 327 } 328 jj = jj + 1 329 } 330 ii = ii + 1 331 } 332 ww = ww + 1 333 } 334 hh = hh + 1 335 } 336 337 let n_x0: *u8 = "x0_latent" as *u8 338 let c_x0: i64 = nx_genfix_dims(M, n_x0, br_strlen(n_x0), ne) 339 if c_x0 >= 0 { 340 let x0: *u8 = nx_genfix_load(M, n_x0, br_strlen(n_x0), c_x0) 341 if (x0 as i64) != 0 { 342 let third: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(3)) 343 let three: i64 = nx_i32_to_f32(3) 344 let taeh: *u8 = sys_mmap(LC * LH * LW * 4 + 64) 345 var q: i64 = 0 346 while q < LC * LH * LW { 347 let zz: i64 = __f32_add(nx_le_read_u32(x0, q * 4), 348 nx_le_read_u32(lat, q * 4) ^ 0x80000000) 349 nx_le_write_u32(taeh, q * 4, __f32_mul(nx_f32_tanh(__f32_mul(zz, third)), three)) 350 q = q + 1 351 } 352 let fd: i64 = sys_openat_wr("/mnt/c/Users/elder/nishi-fixtures/zimage_turbo_2602_q8/tae_h_sov.f32" as *u8, 0x1a4) 353 if fd >= 0 { 354 sys_write(fd, taeh, LC * LH * LW * 4) 355 sys_close(fd) 356 nx_genver_emit("wrote_tae_h_sov" as *u8, LC * LH * LW) 357 } 358 } 359 } 360 361 // ---- grade ------------------------------------------------------------------------- 362 var n_ref: *u8 = "final_out" as *u8 363 if argc >= 5 { n_ref = argv[4] as *u8 } 364 let c_ref: i64 = nx_genfix_dims(M, n_ref, br_strlen(n_ref), ne) 365 if c_ref < 0 { df_puts("missing reference\n" as *u8); return 60 } 366 let ref: *u8 = nx_genfix_load(M, n_ref, br_strlen(n_ref), c_ref) 367 if (ref as i64) == 0 { df_puts("load reference failed\n" as *u8); return 61 } 368 369 let tol: *i64 = sys_mmap(64) as *i64 370 nx_genver_tols(tol) 371 let c: *i64 = sys_mmap(128) as *i64 372 nx_genver_init(c, 0) 373 var f: i64 = 0 374 while f < T * OUTD { 375 nx_genver_tally(c, tol, nx_le_read_u32(fout, f * 4), nx_le_read_u32(ref, f * 4), f) 376 f = f + 1 377 } 378 return nx_genver_report(c) 379}