code wiki / (root) / nx_gen_taedec.nx

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}