code wiki / (root) / nx_gen_ditchain.nx

nx_gen_ditchain.nx source

↩ module page · 192 lines · 8276 B

1// nx_gen_ditchain.nx -- chain N sovereign DiT blocks, weights read straight from the GGUF. 2// 3// This is the sovereign DiT forward stack: given the layer-0 input, the timestep embedding and 4// the rope table, it runs N transformer blocks back to back with NO oracle tensor of any kind in 5// the loop -- every weight comes from the model file, every adaLN vector is computed, every 6// activation is the previous block's own output. 7// 8// Usage: nx_gen_ditchain <model_id> <gguf_path> <n_layers> [n_workers] [reference_name] 9// 10// The reference defaults to `layers.<n_layers>.blk_in`, i.e. the oracle's input to the block 11// AFTER the ones we ran -- so a 2-layer run is graded against layers.2.blk_in. 12// 13// WHAT THIS ANSWERS THAT nx_gen_blockrun DOES NOT 14// blockrun proves ONE assembled block is correct. Chaining asks whether the error stays bounded 15// as blocks compose: 36 layers is 36x the opportunity for a small bias to accumulate into a 16// visibly different image. Per-op fidelity says nothing about that, and neither does one block. 17// 18// ⚠The pass band defaults to 1e-1 because the reference is the ORACLE, which carries its own 19// activation-quantization error at every matmul of every layer it ran. Grading a MORE accurate 20// implementation tightly against it would fail correct work. To ask "did OUR chain drift", grade 21// against an exact reference instead. 22// license_tier: ORIGINAL 23 24import "nx_syscalls.nx" 25import "nx_le.nx" 26import "nx_f32.nx" 27import "nx_f32_div.nx" 28import "nx_f32_cvt.nx" 29import "nx_f16.nx" 30import "nx_f32_exp.nx" 31import "nx_f32_activations.nx" 32import "nx_strconv.nx" 33import "nx_genfix.nx" 34import "nx_genver.nx" 35import "nx_genweights.nx" 36import "nx_genblock.nx" 37 38func dc_puts(s: *u8) -> i64 { 39 var n: i64 = 0 40 while s[n] != (0 as u8) { n = n + 1 } 41 return sys_write(1, s, n) 42} 43 44func main(argc: i64, argv: *i64) -> i64 { 45 if argc < 4 { 46 dc_puts("usage: nx_gen_ditchain <model_id> <gguf_path> <n_layers> [workers] [reference]\n" as *u8) 47 return 2 48 } 49 let M: *u8 = argv[1] as *u8 50 let GP: *u8 = argv[2] as *u8 51 let ep: *i64 = sys_mmap(32) as *i64 52 ep[0] = 0 53 let NL: i64 = nx_strconv_parse_i64(argv[3] as *u8, ep) 54 if ep[0] != 0 { dc_puts("bad n_layers\n" as *u8); return 3 } 55 if NL < 1 { dc_puts("bad n_layers\n" as *u8); return 4 } 56 var NW: i64 = 16 57 if argc >= 5 { 58 ep[0] = 0 59 NW = nx_strconv_parse_i64(argv[4] as *u8, ep) 60 if ep[0] != 0 { NW = 16 } 61 if NW < 1 { NW = 1 } 62 } 63 64 let gw: *i64 = nx_gw_open(GP) 65 if (gw as i64) == 0 { dc_puts("gguf open/parse failed\n" as *u8); return 20 } 66 // Bind the geometry from THIS checkpoint's own tensor shapes before any block runs. 67 // A hardcoded head count is a model welded into the binary; probing makes it swappable. 68 if br_arch_bind(gw) != 0 { dc_puts("architecture not derivable from this checkpoint -- refusing 69" as *u8); return 21 } 70 nx_genver_emit("gguf_tensors" as *u8, nx_gw_ntensors(gw)) 71 72 let ne: *i64 = sys_mmap(64) as *i64 73 let n_in: *u8 = "layers.0.blk_in" as *u8 74 let c_in: i64 = nx_genfix_dims(M, n_in, br_strlen(n_in), ne) 75 if c_in < 0 { dc_puts("missing layers.0.blk_in\n" as *u8); return 30 } 76 let D: i64 = ne[0] 77 let T: i64 = ne[1] 78 let QKV: i64 = (BR_N_HEADS * 3) * BR_HEAD_DIM 79 80 let nm0: *u8 = sys_mmap(256) 81 let wi: i64 = nx_gw_find(gw, br_name(nm0, 0, "feed_forward.w1.weight" as *u8), br_strlen(nm0)) 82 if wi < 0 { dc_puts("missing w1 in gguf\n" as *u8); return 31 } 83 let FD: i64 = nx_gw_dim1(gw, wi) 84 85 nx_genver_emit("dim" as *u8, D) 86 nx_genver_emit("tokens" as *u8, T) 87 nx_genver_emit("ffn_dim" as *u8, FD) 88 nx_genver_emit("layers" as *u8, NL) 89 nx_genver_emit("workers" as *u8, NW) 90 91 let blk_in: *u8 = nx_genfix_load(M, n_in, br_strlen(n_in), c_in) 92 if (blk_in as i64) == 0 { dc_puts("load blk_in failed\n" as *u8); return 40 } 93 let n_te: *u8 = "t_emb" as *u8 94 let t_emb: *u8 = nx_genfix_load(M, n_te, br_strlen(n_te), 256) 95 if (t_emb as i64) == 0 { dc_puts("load t_emb failed\n" as *u8); return 41 } 96 let n_pe: *u8 = "pe" as *u8 97 let pe: *u8 = nx_genfix_load(M, n_pe, br_strlen(n_pe), 2 * 2 * (BR_HEAD_DIM / 2) * T) 98 if (pe as i64) == 0 { dc_puts("load pe failed\n" as *u8); return 42 } 99 100 // ---- scratch, allocated ONCE and reused by every layer ------------------------------ 101 // Re-mapping ~200MB per layer would dominate a 36-layer chain and would also hide the real 102 // per-layer cost behind allocator time. 103 let scr: *i64 = sys_mmap(BR_SCR_SLOTS * 8 + 64) as *i64 104 scr[0] = sys_mmap_shared(T * D * 4) as i64 105 scr[1] = sys_mmap_shared(T * QKV * 4) as i64 106 scr[2] = sys_mmap_shared(T * BR_N_HEADS * BR_HEAD_DIM * 4) as i64 107 scr[3] = sys_mmap_shared(T * BR_N_HEADS * BR_HEAD_DIM * 4) as i64 108 scr[4] = sys_mmap_shared(T * D * 4) as i64 109 scr[5] = sys_mmap_shared(T * D * 4) as i64 110 scr[6] = sys_mmap_shared(T * D * 4) as i64 111 scr[7] = sys_mmap_shared(T * D * 4) as i64 112 scr[8] = sys_mmap_shared(T * D * 4) as i64 113 scr[9] = sys_mmap_shared(T * FD * 4) as i64 114 scr[10] = sys_mmap_shared(T * FD * 4) as i64 115 scr[11] = sys_mmap_shared(T * FD * 4) as i64 116 scr[12] = sys_mmap_shared(T * D * 4) as i64 117 scr[13] = sys_mmap_shared(T * D * 4) as i64 118 scr[14] = sys_mmap((D / 32) * QKV * 8) as i64 119 scr[15] = sys_mmap((D / 32) * D * 8) as i64 120 scr[16] = sys_mmap((D / 32) * FD * 8) as i64 121 scr[17] = sys_mmap((D / 32) * FD * 8) as i64 122 scr[18] = sys_mmap((FD / 32) * D * 8) as i64 123 scr[19] = sys_mmap(BR_CHUNKS * D * 4 + 64) as i64 124 scr[20] = sys_mmap(D * 8) as i64 125 scr[21] = sys_mmap(D * 8) as i64 126 scr[22] = sys_mmap(D * 8) as i64 127 scr[23] = sys_mmap(D * 8) as i64 128 129 // Two full-size buffers, ping-ponged: block L reads `cur` and writes `nxt`. 130 let bufA: *u8 = sys_mmap_shared(T * D * 4) 131 let bufB: *u8 = sys_mmap_shared(T * D * 4) 132 var cur: *u8 = blk_in 133 var nxt: *u8 = bufA 134 135 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(1000000)) 136 let t0: i64 = sys_now_us() 137 138 var lay: i64 = 0 139 while lay < NL { 140 let rc: i64 = br_block(gw, lay, cur, nxt, scr, T, D, QKV, FD, NW, pe, t_emb, eps) 141 if rc != 0 { 142 nx_genver_emit("block_failed_layer" as *u8, lay) 143 nx_genver_emit("rc" as *u8, rc) 144 return 50 145 } 146 nx_genver_emit("done_layer" as *u8, lay) 147 cur = nxt 148 if (nxt as i64) == (bufA as i64) { nxt = bufB } else { nxt = bufA } 149 lay = lay + 1 150 } 151 let t1: i64 = sys_now_us() 152 nx_genver_emit("chain_us" as *u8, t1 - t0) 153 nx_genver_emit("us_per_layer" as *u8, (t1 - t0) / NL) 154 155 // ---- grade ------------------------------------------------------------------------- 156 let rn: *u8 = sys_mmap(256) 157 var n_ref: *u8 = 0 as *u8 158 if argc >= 6 { 159 n_ref = argv[5] as *u8 160 } else { 161 // default: the oracle's input to the block after the last one we ran 162 let pre: *u8 = "layers." as *u8 163 var o: i64 = 0 164 var i: i64 = 0 165 while pre[i] != (0 as u8) { rn[o] = pre[i]; o = o + 1; i = i + 1 } 166 let dec: *u8 = sys_mmap(32) 167 let nd: i64 = nx_strconv_format_i64(NL, dec) 168 var k: i64 = 0 169 while k < nd { rn[o] = dec[k]; o = o + 1; k = k + 1 } 170 let suf: *u8 = ".blk_in" as *u8 171 i = 0 172 while suf[i] != (0 as u8) { rn[o] = suf[i]; o = o + 1; i = i + 1 } 173 rn[o] = 0 174 n_ref = rn 175 } 176 177 let c_ref: i64 = nx_genfix_dims(M, n_ref, br_strlen(n_ref), ne) 178 if c_ref < 0 { dc_puts("missing reference tensor\n" as *u8); return 60 } 179 let ref: *u8 = nx_genfix_load(M, n_ref, br_strlen(n_ref), c_ref) 180 if (ref as i64) == 0 { dc_puts("load reference failed\n" as *u8); return 61 } 181 182 let tol: *i64 = sys_mmap(64) as *i64 183 nx_genver_tols(tol) 184 let c: *i64 = sys_mmap(128) as *i64 185 nx_genver_init(c, 0) 186 var f: i64 = 0 187 while f < T * D { 188 nx_genver_tally(c, tol, nx_le_read_u32(cur, f * 4), nx_le_read_u32(ref, f * 4), f) 189 f = f + 1 190 } 191 return nx_genver_report(c) 192}