nx_gen_taedec.nx source
↩ module page · 360 lines · 13583 B
1// nx_gen_taedec.nx -- SOVEREIGN TAESD (tiny VAE) decoder: latent -> RGB pixels.
2//
3// The last stage before an image. TAESD is 9.4MB of plain f32 3x3 convolutions where the full VAE
4// is 320MB with attention, so it is the cheapest honest route from a sovereign latent to a
5// sovereign picture.
6//
7// h = tanh(z/3)*3 (the caller's fixture already starts here)
8// conv3x3 16->64 (+bias), relu
9// [TAEBlock x3] upsample2 conv3x3 64->64 (no bias)
10// [TAEBlock x3] upsample2 conv3x3 64->64 (no bias)
11// [TAEBlock x3] upsample2 conv3x3 64->64 (no bias)
12// [TAEBlock x1] conv3x3 64->3 (+bias)
13// TAEBlock: conv0 -> relu -> conv2 -> relu -> conv4, then + input, then relu
14//
15// Usage: nx_gen_taedec <model_id> <taesd.safetensors> [workers] [in] [ref]
16//
17// ⚠WHY THIS DOES NOT CALL nx_f32_conv2d_forward
18// That organ's inner loop uses nx_f32_add / nx_f32_mul -- the SOFTWARE IEEE-754 pair -- not the
19// hardware __f32_* intrinsics beside them. At this decoder's ~20 G MACs that is the difference
20// between ~25 minutes and seconds. Same defect class this lane has now hit four times, here in a
21// pre-existing shared organ (filed separately; fixing it there is a fleet change, not a local one).
22// ★★★★★ WHEN A SOVEREIGN KERNEL IS 10x UNDER PEAK, SUSPECT AN nx_-PREFIXED HELPER.
23//
24// Layout is planar packed f32 [C][H][W], W contiguous -- the same order ggml uses ([W,H,C,N]) and
25// the same order safetensors stores conv weights ([C_out][C_in][KH][KW]).
26// license_tier: ORIGINAL
27
28import "nx_syscalls.nx"
29import "nx_le.nx"
30import "nx_f32.nx"
31import "nx_f32_div.nx"
32import "nx_f32_cvt.nx"
33import "nx_f16.nx"
34import "nx_strconv.nx"
35import "nx_genfix.nx"
36import "nx_genver.nx"
37import "nx_safetensors_load.nx"
38
39const TD_CH: i64 = 64
40
41func td_puts(s: *u8) -> i64 {
42 var n: i64 = 0
43 while s[n] != (0 as u8) { n = n + 1 }
44 return sys_write(1, s, n)
45}
46func td_strlen(s: *u8) -> i64 {
47 var n: i64 = 0
48 while s[n] != (0 as u8) { n = n + 1 }
49 return n
50}
51
52// ---- conv 3x3, stride 1, pad 1, over a band of output channels -------------------------
53func td_conv_band(inp: *u8, out: *u8, w: *u8, bias: *u8,
54 C_in: i64, C_out: i64, H: i64, W: i64, co0: i64, co1: i64) -> i64 {
55 let plane: i64 = H * W
56 var co: i64 = co0
57 while co < co1 {
58 var bv: i64 = 0
59 if (bias as i64) != 0 { bv = nx_le_read_u32(bias, co * 4) }
60 var oh: i64 = 0
61 while oh < H {
62 var ow: i64 = 0
63 while ow < W {
64 var acc: i64 = bv
65 var ci: i64 = 0
66 while ci < C_in {
67 let ib: i64 = ci * plane
68 let wb: i64 = (co * C_in + ci) * 9
69 var kh: i64 = 0
70 while kh < 3 {
71 let ih: i64 = oh + kh - 1
72 if ih >= 0 {
73 if ih < H {
74 var kw: i64 = 0
75 while kw < 3 {
76 let iw: i64 = ow + kw - 1
77 if iw >= 0 {
78 if iw < W {
79 acc = __f32_add(acc,
80 __f32_mul(nx_le_read_u32(inp, (ib + ih * W + iw) * 4),
81 nx_le_read_u32(w, (wb + kh * 3 + kw) * 4)))
82 }
83 }
84 kw = kw + 1
85 }
86 }
87 }
88 kh = kh + 1
89 }
90 ci = ci + 1
91 }
92 nx_le_write_u32(out, (co * plane + oh * W + ow) * 4, acc)
93 ow = ow + 1
94 }
95 oh = oh + 1
96 }
97 co = co + 1
98 }
99 return 0
100}
101
102func td_conv(inp: *u8, out: *u8, w: *u8, bias: *u8,
103 C_in: i64, C_out: i64, H: i64, W: i64, nw: i64) -> i64 {
104 if nw <= 1 { td_conv_band(inp, out, w, bias, C_in, C_out, H, W, 0, C_out); return 0 }
105 let pids: *i64 = sys_mmap(nw * 8 + 64) as *i64
106 var k: i64 = 0
107 while k < nw {
108 let a: i64 = k * C_out / nw
109 let b: i64 = (k + 1) * C_out / nw
110 let pid: i64 = sys_fork()
111 if pid == 0 { td_conv_band(inp, out, w, bias, C_in, C_out, H, W, a, b); sys_exit(0) }
112 pids[k] = pid
113 k = k + 1
114 }
115 let st: *i64 = sys_mmap(64) as *i64
116 k = 0
117 while k < nw { sys_wait4(pids[k], st, 0); k = k + 1 }
118 return 0
119}
120
121func td_relu(x: *u8, n: i64) -> i64 {
122 var i: i64 = 0
123 while i < n {
124 let v: i64 = nx_le_read_u32(x, i * 4)
125 if (v & 0x80000000) != 0 { nx_le_write_u32(x, i * 4, 0) }
126 i = i + 1
127 }
128 return 0
129}
130
131func td_add(a: *u8, b: *u8, n: i64) -> i64 {
132 var i: i64 = 0
133 while i < n {
134 nx_le_write_u32(a, i * 4, __f32_add(nx_le_read_u32(a, i * 4), nx_le_read_u32(b, i * 4)))
135 i = i + 1
136 }
137 return 0
138}
139
140func td_copy(dst: *u8, src: *u8, n: i64) -> i64 {
141 var i: i64 = 0
142 while i < n { nx_le_write_u32(dst, i * 4, nx_le_read_u32(src, i * 4)); i = i + 1 }
143 return 0
144}
145
146// nearest-neighbour 2x upscale, planar
147func td_up2(inp: *u8, out: *u8, C: i64, H: i64, W: i64) -> i64 {
148 let OH: i64 = H * 2
149 let OW: i64 = W * 2
150 var c: i64 = 0
151 while c < C {
152 var oh: i64 = 0
153 while oh < OH {
154 var ow: i64 = 0
155 while ow < OW {
156 let v: i64 = nx_le_read_u32(inp, (c * H * W + (oh / 2) * W + (ow / 2)) * 4)
157 nx_le_write_u32(out, (c * OH * OW + oh * OW + ow) * 4, v)
158 ow = ow + 1
159 }
160 oh = oh + 1
161 }
162 c = c + 1
163 }
164 return 0
165}
166
167// ---- safetensors weight -> packed f32 ---------------------------------------------------
168func td_load(buf: *u8, hlen: i64, dstart: i64, name: *u8, n: i64) -> *u8 {
169 let dt: *u8 = sys_mmap(32)
170 let offs: *i64 = sys_mmap(64) as *i64
171 if stl_tensor(buf, hlen, name, dt, offs) != 1 { return 0 as *u8 }
172 let tmp: *i64 = sys_mmap(n * 8 + 64) as *i64
173 // stl_read_f32 returns the ELEMENT COUNT, not a boolean -- a positive rc is success. Checking
174 // `!= 1` treated every real tensor as a failure. The estate banked this exact shape before
175 // (sts_seed returns a row count). ★ A POSITIVE RETURN IS NOT ALWAYS A BOOLEAN.
176 // Requiring the EXACT expected count also catches a tensor whose shape is not what we assumed.
177 let got: i64 = stl_read_f32(buf, dstart, offs, dt, tmp)
178 if got != n { return 0 as *u8 }
179 let out: *u8 = sys_mmap(n * 4 + 64)
180 var i: i64 = 0
181 while i < n { nx_le_write_u32(out, i * 4, tmp[i]); i = i + 1 }
182 return out
183}
184
185// build "decoder.layers.<n>.<suffix>"
186func td_name(out: *u8, layer: i64, suffix: *u8) -> *u8 {
187 let pre: *u8 = "decoder.layers." as *u8
188 var o: i64 = 0
189 var i: i64 = 0
190 while pre[i] != (0 as u8) { out[o] = pre[i]; o = o + 1; i = i + 1 }
191 let dec: *u8 = sys_mmap(32)
192 let nd: i64 = nx_strconv_format_i64(layer, dec)
193 var k: i64 = 0
194 while k < nd { out[o] = dec[k]; o = o + 1; k = k + 1 }
195 out[o] = 0x2E; o = o + 1
196 i = 0
197 while suffix[i] != (0 as u8) { out[o] = suffix[i]; o = o + 1; i = i + 1 }
198 out[o] = 0
199 return out
200}
201
202static g_buf: i64
203static g_hlen: i64
204static g_dstart: i64
205static g_nw: i64
206
207// one plain conv layer, in-place ping-pong
208func td_layer_conv(a: *u8, b: *u8, layer: i64, C_in: i64, C_out: i64, H: i64, W: i64, has_bias: i64) -> i64 {
209 let nm: *u8 = sys_mmap(256)
210 let w: *u8 = td_load(g_buf as *u8, g_hlen, g_dstart, td_name(nm, layer, "weight" as *u8), C_out * C_in * 9)
211 if (w as i64) == 0 { return 0 - 1 }
212 var bias: *u8 = 0 as *u8
213 if has_bias != 0 {
214 bias = td_load(g_buf as *u8, g_hlen, g_dstart, td_name(nm, layer, "bias" as *u8), C_out)
215 if (bias as i64) == 0 { return 0 - 2 }
216 }
217 td_conv(a, b, w, bias, C_in, C_out, H, W, g_nw)
218 return 0
219}
220
221// TAEBlock: conv0 -> relu -> conv2 -> relu -> conv4, + input, relu
222func td_block(x: *u8, t1: *u8, t2: *u8, layer: i64, C: i64, H: i64, W: i64) -> i64 {
223 let nm: *u8 = sys_mmap(256)
224 let n: i64 = C * H * W
225 let w0: *u8 = td_load(g_buf as *u8, g_hlen, g_dstart, td_name(nm, layer, "conv.0.weight" as *u8), C * C * 9)
226 if (w0 as i64) == 0 { return 0 - 1 }
227 let b0: *u8 = td_load(g_buf as *u8, g_hlen, g_dstart, td_name(nm, layer, "conv.0.bias" as *u8), C)
228 let w2: *u8 = td_load(g_buf as *u8, g_hlen, g_dstart, td_name(nm, layer, "conv.2.weight" as *u8), C * C * 9)
229 if (w2 as i64) == 0 { return 0 - 2 }
230 let b2: *u8 = td_load(g_buf as *u8, g_hlen, g_dstart, td_name(nm, layer, "conv.2.bias" as *u8), C)
231 let w4: *u8 = td_load(g_buf as *u8, g_hlen, g_dstart, td_name(nm, layer, "conv.4.weight" as *u8), C * C * 9)
232 if (w4 as i64) == 0 { return 0 - 3 }
233 let b4: *u8 = td_load(g_buf as *u8, g_hlen, g_dstart, td_name(nm, layer, "conv.4.bias" as *u8), C)
234
235 td_conv(x, t1, w0, b0, C, C, H, W, g_nw)
236 td_relu(t1, n)
237 td_conv(t1, t2, w2, b2, C, C, H, W, g_nw)
238 td_relu(t2, n)
239 td_conv(t2, t1, w4, b4, C, C, H, W, g_nw)
240 td_add(t1, x, n) // residual: n_in == n_out here, so the skip is identity
241 td_relu(t1, n)
242 td_copy(x, t1, n)
243 return 0
244}
245
246func main(argc: i64, argv: *i64) -> i64 {
247 if argc < 3 {
248 td_puts("usage: nx_gen_taedec <model_id> <taesd.safetensors> [workers] [in] [ref]\n" as *u8)
249 return 2
250 }
251 let M: *u8 = argv[1] as *u8
252 let SP: *u8 = argv[2] as *u8
253 g_nw = 16
254 if argc >= 4 {
255 let ep: *i64 = sys_mmap(32) as *i64
256 ep[0] = 0
257 g_nw = nx_strconv_parse_i64(argv[3] as *u8, ep)
258 if ep[0] != 0 { g_nw = 16 }
259 if g_nw < 1 { g_nw = 1 }
260 }
261 var in_name: *u8 = "tae_h" as *u8
262 if argc >= 5 { in_name = argv[4] as *u8 }
263 var ref_name: *u8 = "tae_out" as *u8
264 if argc >= 6 { ref_name = argv[5] as *u8 }
265
266 let lenp: *i64 = sys_mmap(32) as *i64
267 lenp[0] = 0
268 let sbuf: *u8 = sys_map_file(SP, lenp)
269 if (sbuf as i64) == 0 { td_puts("safetensors map failed\n" as *u8); return 20 }
270 g_buf = sbuf as i64
271 g_hlen = stl_header_len(sbuf)
272 g_dstart = stl_data_start(sbuf)
273 nx_genver_emit("safetensors_bytes" as *u8, lenp[0])
274
275 let ne: *i64 = sys_mmap(64) as *i64
276 let c_in: i64 = nx_genfix_dims(M, in_name, td_strlen(in_name), ne)
277 if c_in < 0 { td_puts("missing input fixture\n" as *u8); return 30 }
278 let W0: i64 = ne[0]
279 let H0: i64 = ne[1]
280 let ZC: i64 = ne[2]
281 nx_genver_emit("latent_w" as *u8, W0)
282 nx_genver_emit("latent_h" as *u8, H0)
283 nx_genver_emit("latent_c" as *u8, ZC)
284 nx_genver_emit("workers" as *u8, g_nw)
285
286 let z: *u8 = nx_genfix_load(M, in_name, td_strlen(in_name), c_in)
287 if (z as i64) == 0 { td_puts("load input failed\n" as *u8); return 31 }
288
289 // Buffers sized for the LARGEST stage (8x upscale, 64 channels).
290 let FW: i64 = W0 * 8
291 let FH: i64 = H0 * 8
292 let big: i64 = TD_CH * FH * FW * 4 + 64
293 let a: *u8 = sys_mmap_shared(big)
294 let b: *u8 = sys_mmap_shared(big)
295 let c: *u8 = sys_mmap_shared(big)
296
297 let t0: i64 = sys_now_us()
298
299 // layer 0: conv 16->64 (+bias), then relu
300 if td_layer_conv(z, a, 0, ZC, TD_CH, H0, W0, 1) != 0 { td_puts("layer0 failed\n" as *u8); return 40 }
301 td_relu(a, TD_CH * H0 * W0)
302
303 var H: i64 = H0
304 var W: i64 = W0
305 // three stages of [3 blocks -> upsample -> conv], then a final block + output conv
306 let stage_first: *i64 = sys_mmap(64) as *i64
307 stage_first[0] = 2
308 stage_first[1] = 7
309 stage_first[2] = 12
310 let stage_conv: *i64 = sys_mmap(64) as *i64
311 stage_conv[0] = 6
312 stage_conv[1] = 11
313 stage_conv[2] = 16
314 var st: i64 = 0
315 while st < 3 {
316 var bi: i64 = 0
317 while bi < 3 {
318 if td_block(a, b, c, stage_first[st] + bi, TD_CH, H, W) != 0 {
319 nx_genver_emit("block_failed" as *u8, stage_first[st] + bi)
320 return 41
321 }
322 bi = bi + 1
323 }
324 td_up2(a, b, TD_CH, H, W)
325 H = H * 2
326 W = W * 2
327 if td_layer_conv(b, a, stage_conv[st], TD_CH, TD_CH, H, W, 0) != 0 {
328 nx_genver_emit("stage_conv_failed" as *u8, stage_conv[st])
329 return 42
330 }
331 st = st + 1
332 }
333 if td_block(a, b, c, 17, TD_CH, H, W) != 0 { td_puts("block 17 failed\n" as *u8); return 43 }
334 if td_layer_conv(a, b, 18, TD_CH, 3, H, W, 1) != 0 { td_puts("layer18 failed\n" as *u8); return 44 }
335
336 let t1: i64 = sys_now_us()
337 nx_genver_emit("decode_us" as *u8, t1 - t0)
338 nx_genver_emit("out_w" as *u8, W)
339 nx_genver_emit("out_h" as *u8, H)
340
341 let c_ref: i64 = nx_genfix_dims(M, ref_name, td_strlen(ref_name), ne)
342 if c_ref < 0 { td_puts("missing reference\n" as *u8); return 60 }
343 let ref: *u8 = nx_genfix_load(M, ref_name, td_strlen(ref_name), c_ref)
344 if (ref as i64) == 0 { td_puts("load reference failed\n" as *u8); return 61 }
345
346 // Also write the decoded planes raw so a PNG can be made from OUR pixels, not the oracle's.
347 let ofd: i64 = sys_openat_wr("/tmp/nx_taedec_rgb.f32" as *u8, 0x1a4)
348 if ofd >= 0 { sys_write(ofd, b, 3 * H * W * 4); sys_close(ofd) }
349
350 let tol: *i64 = sys_mmap(64) as *i64
351 nx_genver_tols(tol)
352 let cc: *i64 = sys_mmap(128) as *i64
353 nx_genver_init(cc, 1)
354 var f: i64 = 0
355 while f < 3 * H * W {
356 nx_genver_tally(cc, tol, nx_le_read_u32(b, f * 4), nx_le_read_u32(ref, f * 4), f)
357 f = f + 1
358 }
359 return nx_genver_report(cc)
360}