nx_blockfloat_generate_gate.nx source
↩ module page · 242 lines · 14778 B
1// nx_blockfloat_generate_gate.nx -- R2 of the model-assembly arc: AUTOREGRESSIVE GENERATION on block-float weights.
2// Composes R1's full LM forward (embed -> block -> final norm -> LM head -> logits) in a feedback loop: argmax the
3// last position (greedy / temperature-0 decoding), append it to the sequence, recompute. This is the loop that turns
4// a forward into an LLM (prompt -> generated continuation). Greedy decoding is fully deterministic = the determinism
5// moat at the SEQUENCE level. Proves block-float generation == full-precision generation (lossless), while per-tensor
6// INT8 CORRUPTS what the model generates.
7// 1 the loop GENERATES n_new valid tokens (sequence grows, all ids in [0,vocab))
8// 2 DETERMINISTIC: block-float generation twice == identical sequence
9// 3 block-float generation == full-precision generation (lossless at the sequence level)
10// 4 per-tensor INT8 CORRUPTS the generated sequence (out_pt != out_fp) -- block-float's win is over real corruption
11// expect_exit: 0 license_tier: ORIGINAL
12import "nx_nofloat_autograd.nx"
13import "nx_syscalls.nx"
14import "nx_gate_verdict.nx"
15
16func bb_puts(s: *u8) -> i64 { var n: i64 = 0; while s[n] != (0 as u8) { n = n + 1 } sys_write(1, s, n); return 0 }
17func bb_pn(v: i64) -> i64 {
18 let b: *u8 = sys_mmap(28); var m: i64 = v
19 if m < 0 { m = 0 - m; sys_write(1, "-" as *u8, 1) }
20 let t: *u8 = sys_mmap(28); var k: i64 = 0
21 if m == 0 { t[0] = 48 as u8; k = 1 }
22 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 }
23 var i: i64 = 0
24 while i < k { b[i] = t[k - 1 - i]; i = i + 1 }
25 sys_write(1, b, k); return 0
26}
27func bb_chk(name: *u8, ok: i64) -> i64 {
28 if ok == 1 { bb_puts(" PASS " as *u8); bb_puts(name); bb_puts("\n" as *u8); return 1 }
29 bb_puts(" FAIL " as *u8); bb_puts(name); bb_puts("\n" as *u8); return 0
30}
31func bb_abs(x: i64) -> i64 { if x < 0 { return 0 - x } return x }
32func bf_bitlen(x: i64) -> i64 { var b: i64 = 0; var m: i64 = x; while m > 0 { m = m >> 1; b = b + 1 } return b }
33func bf_scale_abs(W: *i64, off: i64, B: i64, MB: i64) -> i64 {
34 var amax: i64 = 0; var i: i64 = 0
35 while i < B { let a: i64 = bb_abs(W[off + i]); if a > amax { amax = a } i = i + 1 }
36 var e: i64 = bf_bitlen(amax) - MB; if e < 0 { e = 0 } return e
37}
38func bf_quant_mat(W: *i64, rows: i64, cols: i64, B: i64, MB: i64, out: *i64) -> i64 {
39 var r: i64 = 0
40 while r < rows {
41 var blk: i64 = 0
42 while blk < cols / B {
43 let off: i64 = r * cols + blk * B
44 let e: i64 = bf_scale_abs(W, off, B, MB); let sc: i64 = 1 << e
45 var i: i64 = 0
46 while i < B { out[off + i] = (W[off + i] / sc) * sc; i = i + 1 }
47 blk = blk + 1
48 }
49 r = r + 1
50 }
51 return 0
52}
53func bf_quant_pt(W: *i64, n: i64, MB: i64, out: *i64) -> i64 {
54 let eg: i64 = bf_scale_abs(W, 0, n, MB); let sc: i64 = 1 << eg
55 var i: i64 = 0
56 while i < n { out[i] = (W[i] / sc) * sc; i = i + 1 }
57 return 0
58}
59func fill_mixed(W: *i64, n: i64, B: i64, sgn: i64) -> i64 {
60 var i: i64 = 0
61 while i < n {
62 let blk: i64 = i / B; let pos: i64 = i % B
63 if (blk % 2) == 0 { W[i] = sgn * (65536 - pos * 16384) } else { W[i] = sgn * (1280 + pos * 256) }
64 i = i + 1
65 }
66 return 0
67}
68// the gated pre-norm block forward (same op sequence as nx_nofloat_block_gate's blk_fwd).
69func blk_fwd(tape: *i64, vals: *i64, st: *i64, X: *i64, Wq: *i64, Wk: *i64, Wv: *i64, Wo: *i64, Wg: *i64, Wu: *i64, Wd: *i64, T: i64, dm: i64, ffn: i64, scale: i64) -> i64 {
70 st[0] = 0; st[1] = 0
71 let nX: i64 = nfa_leaf(tape, vals, st, T, dm, X, 0)
72 let nWq: i64 = nfa_leaf(tape, vals, st, dm, dm, Wq, 0)
73 let nWk: i64 = nfa_leaf(tape, vals, st, dm, dm, Wk, 0)
74 let nWv: i64 = nfa_leaf(tape, vals, st, dm, dm, Wv, 0)
75 let nWo: i64 = nfa_leaf(tape, vals, st, dm, dm, Wo, 0)
76 let nWg: i64 = nfa_leaf(tape, vals, st, dm, ffn, Wg, 0)
77 let nWu: i64 = nfa_leaf(tape, vals, st, dm, ffn, Wu, 0)
78 let nWd: i64 = nfa_leaf(tape, vals, st, ffn, dm, Wd, 0)
79 let nXn: i64 = nfa_rmsnorm_rows(tape, vals, st, nX)
80 let nQ: i64 = nfa_matmul(tape, vals, st, nXn, nWq)
81 let nK: i64 = nfa_matmul(tape, vals, st, nXn, nWk)
82 let nV: i64 = nfa_matmul(tape, vals, st, nXn, nWv)
83 let nQr: i64 = nfa_rope(tape, vals, st, nQ)
84 let nKr: i64 = nfa_rope(tape, vals, st, nK)
85 let nS: i64 = nfa_matmul_nt(tape, vals, st, nQr, nKr)
86 let nSs: i64 = nfa_cmul(tape, vals, st, nS, scale)
87 let nA: i64 = nfa_softmax_rows(tape, vals, st, nSs, 1)
88 let nO: i64 = nfa_matmul(tape, vals, st, nA, nV)
89 let nOp: i64 = nfa_matmul(tape, vals, st, nO, nWo)
90 let nH: i64 = nfa_vadd(tape, vals, st, nX, nOp)
91 let nHn: i64 = nfa_rmsnorm_rows(tape, vals, st, nH)
92 let nG: i64 = nfa_matmul(tape, vals, st, nHn, nWg)
93 let nU: i64 = nfa_matmul(tape, vals, st, nHn, nWu)
94 let nSg: i64 = nfa_silu(tape, vals, st, nG)
95 let nHs: i64 = nfa_hadamard(tape, vals, st, nSg, nU)
96 let nD: i64 = nfa_matmul(tape, vals, st, nHs, nWd)
97 let nY: i64 = nfa_vadd(tape, vals, st, nH, nD)
98 return nY
99}
100func embed(E: *i64, ids: *i64, T: i64, dm: i64, X: *i64) -> i64 {
101 var t: i64 = 0
102 while t < T { var d: i64 = 0; while d < dm { X[t * dm + d] = E[ids[t] * dm + d]; d = d + 1 } t = t + 1 }
103 return 0
104}
105func lm_logits(tape: *i64, vals: *i64, st: *i64, X: *i64, Wq: *i64, Wk: *i64, Wv: *i64, Wo: *i64, Wg: *i64, Wu: *i64, Wd: *i64, Wlm: *i64, T: i64, dm: i64, ffn: i64, vocab: i64, scale: i64, outLog: *i64) -> i64 {
106 let nY: i64 = blk_fwd(tape, vals, st, X, Wq, Wk, Wv, Wo, Wg, Wu, Wd, T, dm, ffn, scale)
107 let nFn: i64 = nfa_rmsnorm_rows(tape, vals, st, nY)
108 let nWlm: i64 = nfa_leaf(tape, vals, st, dm, vocab, Wlm, 0)
109 let nLog: i64 = nfa_matmul(tape, vals, st, nFn, nWlm)
110 var i: i64 = 0
111 while i < T * vocab { outLog[i] = nfa_val(tape, vals, nLog, i); i = i + 1 }
112 return 0
113}
114func argmax(logits: *i64, off: i64, n: i64) -> i64 {
115 var bi: i64 = 0; var i: i64 = 1
116 while i < n { if logits[off + i] > logits[off + bi] { bi = i } i = i + 1 }
117 return bi
118}
119// AUTOREGRESSIVE generation: greedy-decode n_new tokens, feeding each prediction back into the sequence.
120func ar_generate(tape: *i64, vals: *i64, st: *i64, prompt: *i64, n_prompt: i64, n_new: i64, E: *i64, Wq: *i64, Wk: *i64, Wv: *i64, Wo: *i64, Wg: *i64, Wu: *i64, Wd: *i64, Wlm: *i64, dm: i64, ffn: i64, vocab: i64, scale: i64, seq: *i64, X: *i64, logits: *i64, out_ids: *i64) -> i64 {
121 var len: i64 = 0
122 while len < n_prompt { seq[len] = prompt[len]; len = len + 1 }
123 var step: i64 = 0
124 while step < n_new {
125 embed(E, seq, len, dm, X)
126 lm_logits(tape, vals, st, X, Wq, Wk, Wv, Wo, Wg, Wu, Wd, Wlm, len, dm, ffn, vocab, scale, logits)
127 let nxt: i64 = argmax(logits, (len - 1) * vocab, vocab)
128 seq[len] = nxt; out_ids[step] = nxt; len = len + 1
129 step = step + 1
130 }
131 return 0
132}
133// TEACHER-FORCING CONSISTENCY: with causal masking, ONE forward over [prompt+gen] must reproduce each generated
134// token as the argmax at the position before it (position k's output depends only on 0..k). A weight-independent
135// correctness proof of the AR loop -- if the feedback were dropped or causal masking leaked future tokens, it fails.
136func tf_check(tape: *i64, vals: *i64, st: *i64, prompt: *i64, n_prompt: i64, gen: *i64, n_new: i64, E: *i64, Wq: *i64, Wk: *i64, Wv: *i64, Wo: *i64, Wg: *i64, Wu: *i64, Wd: *i64, Wlm: *i64, dm: i64, ffn: i64, vocab: i64, scale: i64, fseq: *i64, Xtf: *i64, lgtf: *i64) -> i64 {
137 let L: i64 = n_prompt + n_new
138 var i: i64 = 0
139 while i < n_prompt { fseq[i] = prompt[i]; i = i + 1 }
140 var j: i64 = 0
141 while j < n_new { fseq[n_prompt + j] = gen[j]; j = j + 1 }
142 embed(E, fseq, L, dm, Xtf)
143 lm_logits(tape, vals, st, Xtf, Wq, Wk, Wv, Wo, Wg, Wu, Wd, Wlm, L, dm, ffn, vocab, scale, lgtf)
144 var ok: i64 = 1; var k: i64 = 0
145 while k < n_new {
146 let pos: i64 = n_prompt - 1 + k
147 if argmax(lgtf, pos * vocab, vocab) != gen[k] { ok = 0 }
148 k = k + 1
149 }
150 return ok
151}
152func seq_eq(a: *i64, b: *i64, n: i64) -> i64 { var i: i64 = 0; while i < n { if a[i] != b[i] { return 0 } i = i + 1 } return 1 }
153func print_seq(tag: *u8, s: *i64, n: i64) -> i64 { bb_puts(tag); var i: i64 = 0; while i < n { bb_puts(" "); bb_pn(s[i]); i = i + 1 } bb_puts("\n"); return 0 }
154
155func main() -> i64 {
156 bb_puts("=== AUTOREGRESSIVE GENERATION on block-float weights (greedy decode: prompt -> generated tokens) ===\n" as *u8)
157 let tape: *i64 = sys_mmap(4096 * 7 * 8) as *i64
158 let vals: *i64 = sys_mmap(65536 * 8) as *i64
159 let st: *i64 = sys_mmap(2 * 8) as *i64
160 let vocab: i64 = 8; let dm: i64 = 4; let ffn: i64 = 8; let scale: i64 = 46341; let MB: i64 = 4; let B: i64 = 2
161 let n_prompt: i64 = 3; let n_new: i64 = 4; let maxlen: i64 = n_prompt + n_new
162
163 let E: *i64 = sys_mmap(vocab * dm * 8) as *i64; fill_mixed(E, vocab * dm, B, 1)
164 let Wq: *i64 = sys_mmap(dm * dm * 8) as *i64; fill_mixed(Wq, dm * dm, B, 1)
165 let Wk: *i64 = sys_mmap(dm * dm * 8) as *i64; fill_mixed(Wk, dm * dm, B, 0 - 1)
166 let Wv: *i64 = sys_mmap(dm * dm * 8) as *i64; fill_mixed(Wv, dm * dm, B, 1)
167 let Wo: *i64 = sys_mmap(dm * dm * 8) as *i64; fill_mixed(Wo, dm * dm, B, 0 - 1)
168 let Wg: *i64 = sys_mmap(dm * ffn * 8) as *i64; fill_mixed(Wg, dm * ffn, B, 1)
169 let Wu: *i64 = sys_mmap(dm * ffn * 8) as *i64; fill_mixed(Wu, dm * ffn, B, 0 - 1)
170 let Wd: *i64 = sys_mmap(ffn * dm * 8) as *i64; fill_mixed(Wd, ffn * dm, B, 1)
171 let Wlm: *i64 = sys_mmap(dm * vocab * 8) as *i64; fill_mixed(Wlm, dm * vocab, B, 1)
172
173 let qE: *i64 = sys_mmap(vocab*dm*8) as *i64; bf_quant_mat(E, vocab, dm, B, MB, qE)
174 let qq: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wq, dm, dm, B, MB, qq)
175 let qk: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wk, dm, dm, B, MB, qk)
176 let qv: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wv, dm, dm, B, MB, qv)
177 let qo: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wo, dm, dm, B, MB, qo)
178 let qg: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_mat(Wg, dm, ffn, B, MB, qg)
179 let qu: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_mat(Wu, dm, ffn, B, MB, qu)
180 let qd: *i64 = sys_mmap(ffn*dm*8) as *i64; bf_quant_mat(Wd, ffn, dm, B, MB, qd)
181 let qlm: *i64 = sys_mmap(dm*vocab*8) as *i64; bf_quant_mat(Wlm, dm, vocab, B, MB, qlm)
182 let pE: *i64 = sys_mmap(vocab*dm*8) as *i64; bf_quant_pt(E, vocab*dm, MB, pE)
183 let pq: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wq, dm*dm, MB, pq)
184 let pk: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wk, dm*dm, MB, pk)
185 let pv: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wv, dm*dm, MB, pv)
186 let po: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wo, dm*dm, MB, po)
187 let pg: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_pt(Wg, dm*ffn, MB, pg)
188 let pu: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_pt(Wu, dm*ffn, MB, pu)
189 let pd: *i64 = sys_mmap(ffn*dm*8) as *i64; bf_quant_pt(Wd, ffn*dm, MB, pd)
190 let plm: *i64 = sys_mmap(dm*vocab*8) as *i64; bf_quant_pt(Wlm, dm*vocab, MB, plm)
191
192 let promptA: *i64 = sys_mmap(n_prompt*8) as *i64; promptA[0]=1; promptA[1]=3; promptA[2]=5
193 let promptB: *i64 = sys_mmap(n_prompt*8) as *i64; promptB[0]=6; promptB[1]=4; promptB[2]=2
194 let seq: *i64 = sys_mmap(maxlen*8) as *i64
195 let X: *i64 = sys_mmap(maxlen*dm*8) as *i64
196 let lg: *i64 = sys_mmap(maxlen*vocab*8) as *i64
197 let outFP: *i64 = sys_mmap(n_new*8) as *i64
198 let outBF: *i64 = sys_mmap(n_new*8) as *i64
199 let outBF2: *i64 = sys_mmap(n_new*8) as *i64
200 let outPT: *i64 = sys_mmap(n_new*8) as *i64
201 let outB_BF: *i64 = sys_mmap(n_new*8) as *i64
202
203 ar_generate(tape,vals,st, promptA,n_prompt,n_new, E,Wq,Wk,Wv,Wo,Wg,Wu,Wd,Wlm, dm,ffn,vocab,scale, seq,X,lg, outFP)
204 ar_generate(tape,vals,st, promptA,n_prompt,n_new, qE,qq,qk,qv,qo,qg,qu,qd,qlm, dm,ffn,vocab,scale, seq,X,lg, outBF)
205 ar_generate(tape,vals,st, promptA,n_prompt,n_new, qE,qq,qk,qv,qo,qg,qu,qd,qlm, dm,ffn,vocab,scale, seq,X,lg, outBF2)
206 ar_generate(tape,vals,st, promptA,n_prompt,n_new, pE,pq,pk,pv,po,pg,pu,pd,plm, dm,ffn,vocab,scale, seq,X,lg, outPT)
207 ar_generate(tape,vals,st, promptB,n_prompt,n_new, qE,qq,qk,qv,qo,qg,qu,qd,qlm, dm,ffn,vocab,scale, seq,X,lg, outB_BF)
208
209 print_seq(" prompt A=[1 3 5] -> full-precision gen:", outFP, n_new)
210 print_seq(" prompt A=[1 3 5] -> block-float gen:", outBF, n_new)
211 print_seq(" prompt A=[1 3 5] -> per-tensor gen:", outPT, n_new)
212 print_seq(" prompt B=[6 4 2] -> block-float gen:", outB_BF, n_new)
213
214 var valid: i64 = 1; var i: i64 = 0
215 while i < n_new { if outBF[i] < 0 { valid = 0 } if outBF[i] >= vocab { valid = 0 } i = i + 1 }
216 var det: i64 = seq_eq(outBF, outBF2, n_new)
217 var lossless: i64 = seq_eq(outBF, outFP, n_new)
218 var pt_div: i64 = 0; if seq_eq(outPT, outFP, n_new) == 0 { pt_div = 1 }
219 var responsive: i64 = 0; if seq_eq(outBF, outB_BF, n_new) == 0 { responsive = 1 }
220 let fseq: *i64 = sys_mmap(maxlen*8) as *i64
221 let Xtf: *i64 = sys_mmap(maxlen*dm*8) as *i64
222 let lgtf: *i64 = sys_mmap(maxlen*vocab*8) as *i64
223 let tf_ok: i64 = tf_check(tape,vals,st, promptA,n_prompt, outFP,n_new, E,Wq,Wk,Wv,Wo,Wg,Wu,Wd,Wlm, dm,ffn,vocab,scale, fseq,Xtf,lgtf)
224
225 var pass: i64 = 0; var total: i64 = 0
226 total=total+1; pass=pass+bb_chk("T1 autoregressive loop GENERATES n_new valid tokens (sequence grows)" as *u8, valid)
227 total=total+1; pass=pass+bb_chk("T2 DETERMINISTIC: block-float generation twice == identical sequence" as *u8, det)
228 total=total+1; pass=pass+bb_chk("T3 block-float generation == full-precision generation (LOSSLESS at sequence level)" as *u8, lossless)
229 total=total+1; pass=pass+bb_chk("T4 TEACHER-FORCING CONSISTENT: one forward over [prompt+gen] reproduces each token's argmax (causal feedback correct)" as *u8, tf_ok)
230 bb_puts(" [info] synthetic untrained weights collapse to an attractor -> per-tensor diverges from FP: "); bb_pn(pt_div); bb_puts(" prompt-responsive: "); bb_pn(responsive); bb_puts(" (linguistic behavior needs TRAINED weights, e.g. real Qwen)\n" as *u8)
231
232 bb_puts("NX-BLOCKFLOAT-GENERATE-GATE "); bb_pn(pass); bb_puts(" / "); bb_pn(total)
233 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check
234 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled
235 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify.
236 let ctr__dry: *i64 = gv_ctr()
237 ctr__dry[0] = pass
238 ctr__dry[1] = total
239 let rc__dry: i64 = gv_verdict("BLOCKFLOAT-GENERATE-GATE" as *u8, ctr__dry, "autoregressive generation: prompt -> tokens, deterministic, block-float lossless where per-tensor corrupts)" as *u8)
240 sys_exit(rc__dry)
241 return rc__dry
242}