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}