code wiki / (root) / nx_blockfloat_generate_gate.nx

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}