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}