code wiki / (root) / nx_blockfloat_stack_gate.nx

nx_blockfloat_stack_gate.nx source

↩ module page · 188 lines · 11207 B

1// nx_blockfloat_stack_gate.nx -- BLOCK-FLOAT weights through a STACKED (multi-block) transformer -- the step 2// toward a full model. Chains the existing gradcheck-verified block forward (blk_fwd) N_LAYERS times (each block's 3// output feeds the next), with the weights block-float quantized. Proves block-float COMPOSES through depth, and 4// the key depth property: per-tensor INT8 degradation COMPOUNDS with depth, while block-float stays stable. 5// 1 the stacked transformer RUNS with block-float weights (non-trivial output) 6// 2 EXCEED: block-float stack error vs full-precision < per-tensor over N_LAYERS blocks 7// 3 DETERMINISTIC: the block-float stack run twice == bit-identical 8// 4 DEPTH: per-tensor error at depth N > per-tensor error at depth 1 (degradation compounds; block-float stable) 9// expect_exit: 0 license_tier: ORIGINAL 10import "nx_nofloat_autograd.nx" 11import "nx_syscalls.nx" 12import "nx_gate_verdict.nx" 13 14func 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 } 15func bb_pn(v: i64) -> i64 { 16 let b: *u8 = sys_mmap(28); var m: i64 = v 17 if m < 0 { m = 0 - m; sys_write(1, "-" as *u8, 1) } 18 let t: *u8 = sys_mmap(28); var k: i64 = 0 19 if m == 0 { t[0] = 48 as u8; k = 1 } 20 while m > 0 { t[k] = (48 + (m % 10)) as u8; m = m / 10; k = k + 1 } 21 var i: i64 = 0 22 while i < k { b[i] = t[k - 1 - i]; i = i + 1 } 23 sys_write(1, b, k); return 0 24} 25func bb_chk(name: *u8, ok: i64) -> i64 { 26 if ok == 1 { bb_puts(" PASS " as *u8); bb_puts(name); bb_puts("\n" as *u8); return 1 } 27 bb_puts(" FAIL " as *u8); bb_puts(name); bb_puts("\n" as *u8); return 0 28} 29func bb_abs(x: i64) -> i64 { if x < 0 { return 0 - x } return x } 30func bf_bitlen(x: i64) -> i64 { var b: i64 = 0; var m: i64 = x; while m > 0 { m = m >> 1; b = b + 1 } return b } 31func bf_scale_abs(W: *i64, off: i64, B: i64, MB: i64) -> i64 { 32 var amax: i64 = 0; var i: i64 = 0 33 while i < B { let a: i64 = bb_abs(W[off + i]); if a > amax { amax = a } i = i + 1 } 34 var e: i64 = bf_bitlen(amax) - MB 35 if e < 0 { e = 0 } 36 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} 59 60func 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, leaves: *i64) -> i64 { 61 st[0] = 0; st[1] = 0 62 let nX: i64 = nfa_leaf(tape, vals, st, T, dm, X, 0) 63 let nWq: i64 = nfa_leaf(tape, vals, st, dm, dm, Wq, 0) 64 let nWk: i64 = nfa_leaf(tape, vals, st, dm, dm, Wk, 0) 65 let nWv: i64 = nfa_leaf(tape, vals, st, dm, dm, Wv, 0) 66 let nWo: i64 = nfa_leaf(tape, vals, st, dm, dm, Wo, 0) 67 let nWg: i64 = nfa_leaf(tape, vals, st, dm, ffn, Wg, 0) 68 let nWu: i64 = nfa_leaf(tape, vals, st, dm, ffn, Wu, 0) 69 let nWd: i64 = nfa_leaf(tape, vals, st, ffn, dm, Wd, 0) 70 let nXn: i64 = nfa_rmsnorm_rows(tape, vals, st, nX) 71 let nQ: i64 = nfa_matmul(tape, vals, st, nXn, nWq) 72 let nK: i64 = nfa_matmul(tape, vals, st, nXn, nWk) 73 let nV: i64 = nfa_matmul(tape, vals, st, nXn, nWv) 74 let nQr: i64 = nfa_rope(tape, vals, st, nQ) 75 let nKr: i64 = nfa_rope(tape, vals, st, nK) 76 let nS: i64 = nfa_matmul_nt(tape, vals, st, nQr, nKr) 77 let nSs: i64 = nfa_cmul(tape, vals, st, nS, scale) 78 let nA: i64 = nfa_softmax_rows(tape, vals, st, nSs, 1) 79 let nO: i64 = nfa_matmul(tape, vals, st, nA, nV) 80 let nOp: i64 = nfa_matmul(tape, vals, st, nO, nWo) 81 let nH: i64 = nfa_vadd(tape, vals, st, nX, nOp) 82 let nHn: i64 = nfa_rmsnorm_rows(tape, vals, st, nH) 83 let nG: i64 = nfa_matmul(tape, vals, st, nHn, nWg) 84 let nU: i64 = nfa_matmul(tape, vals, st, nHn, nWu) 85 let nSg: i64 = nfa_silu(tape, vals, st, nG) 86 let nHs: i64 = nfa_hadamard(tape, vals, st, nSg, nU) 87 let nD: i64 = nfa_matmul(tape, vals, st, nHs, nWd) 88 let nY: i64 = nfa_vadd(tape, vals, st, nH, nD) 89 leaves[0] = nWq; leaves[1] = nWd 90 return nY 91} 92func blk_out(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, outY: *i64) -> i64 { 93 let lv: *i64 = sys_mmap(2 * 8) as *i64 94 let nY: i64 = blk_fwd(tape, vals, st, X, Wq, Wk, Wv, Wo, Wg, Wu, Wd, T, dm, ffn, scale, lv) 95 var i: i64 = 0 96 while i < T * dm { outY[i] = nfa_val(tape, vals, nY, i); i = i + 1 } 97 return 0 98} 99// stack N_LAYERS blocks: each block's output is the next block's input (residuals carry through). 100func stack_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, N_LAYERS: i64, outY: *i64) -> i64 { 101 let cur: *i64 = sys_mmap(T * dm * 8) as *i64 102 let nxt: *i64 = sys_mmap(T * dm * 8) as *i64 103 var i: i64 = 0 104 while i < T * dm { cur[i] = X[i]; i = i + 1 } 105 var L: i64 = 0 106 while L < N_LAYERS { 107 blk_out(tape, vals, st, cur, Wq, Wk, Wv, Wo, Wg, Wu, Wd, T, dm, ffn, scale, nxt) 108 var j: i64 = 0 109 while j < T * dm { cur[j] = nxt[j]; j = j + 1 } 110 L = L + 1 111 } 112 var k: i64 = 0 113 while k < T * dm { outY[k] = cur[k]; k = k + 1 } 114 return 0 115} 116func l1err(a: *i64, b: *i64, n: i64) -> i64 { var e: i64 = 0; var i: i64 = 0; while i < n { e = e + bb_abs(a[i] - b[i]); i = i + 1 } return e } 117 118func main() -> i64 { 119 bb_puts("=== BLOCK-FLOAT weights through a STACKED multi-block transformer (toward a full model) ===\n" as *u8) 120 let tape: *i64 = sys_mmap(1024 * 7 * 8) as *i64 121 let vals: *i64 = sys_mmap(16384 * 8) as *i64 122 let st: *i64 = sys_mmap(2 * 8) as *i64 123 let T: i64 = 2; let dm: i64 = 2; let ffn: i64 = 4; let scale: i64 = 46341; let MB: i64 = 4; let B: i64 = 2; let NL: i64 = 3 124 125 let X: *i64 = sys_mmap(T * dm * 8) as *i64; X[0] = 32768; X[1] = 0 - 16384; X[2] = 49152; X[3] = 24576 126 let Wq: *i64 = sys_mmap(dm * dm * 8) as *i64; Wq[0] = 65536; Wq[1] = 49152; Wq[2] = 1280; Wq[3] = 1536 127 let Wk: *i64 = sys_mmap(dm * dm * 8) as *i64; Wk[0] = 1536; Wk[1] = 1280; Wk[2] = 49152; Wk[3] = 65536 128 let Wv: *i64 = sys_mmap(dm * dm * 8) as *i64; Wv[0] = 49152; Wv[1] = 65536; Wv[2] = 1024; Wv[3] = 1792 129 let Wo: *i64 = sys_mmap(dm * dm * 8) as *i64; Wo[0] = 1280; Wo[1] = 1536; Wo[2] = 65536; Wo[3] = 49152 130 let Wg: *i64 = sys_mmap(dm * ffn * 8) as *i64; Wg[0] = 65536; Wg[1] = 49152; Wg[2] = 1280; Wg[3] = 1536; Wg[4] = 1536; Wg[5] = 1280; Wg[6] = 49152; Wg[7] = 65536 131 let Wu: *i64 = sys_mmap(dm * ffn * 8) as *i64; Wu[0] = 1280; Wu[1] = 1536; Wu[2] = 65536; Wu[3] = 49152; Wu[4] = 49152; Wu[5] = 65536; Wu[6] = 1280; Wu[7] = 1024 132 let Wd: *i64 = sys_mmap(ffn * dm * 8) as *i64; Wd[0] = 65536; Wd[1] = 49152; Wd[2] = 1280; Wd[3] = 1536; Wd[4] = 1536; Wd[5] = 1280; Wd[6] = 49152; Wd[7] = 65536 133 134 let qq: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wq, dm, dm, B, MB, qq) 135 let qk: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wk, dm, dm, B, MB, qk) 136 let qv: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wv, dm, dm, B, MB, qv) 137 let qo: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_mat(Wo, dm, dm, B, MB, qo) 138 let qg: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_mat(Wg, dm, ffn, B, MB, qg) 139 let qu: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_mat(Wu, dm, ffn, B, MB, qu) 140 let qd: *i64 = sys_mmap(ffn*dm*8) as *i64; bf_quant_mat(Wd, ffn, dm, B, MB, qd) 141 let pq: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wq, dm*dm, MB, pq) 142 let pk: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wk, dm*dm, MB, pk) 143 let pv: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wv, dm*dm, MB, pv) 144 let po: *i64 = sys_mmap(dm*dm*8) as *i64; bf_quant_pt(Wo, dm*dm, MB, po) 145 let pg: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_pt(Wg, dm*ffn, MB, pg) 146 let pu: *i64 = sys_mmap(dm*ffn*8) as *i64; bf_quant_pt(Wu, dm*ffn, MB, pu) 147 let pd: *i64 = sys_mmap(ffn*dm*8) as *i64; bf_quant_pt(Wd, ffn*dm, MB, pd) 148 149 // depth-NL stacks 150 let Yr: *i64 = sys_mmap(T*dm*8) as *i64; stack_fwd(tape,vals,st,X,Wq,Wk,Wv,Wo,Wg,Wu,Wd,T,dm,ffn,scale,NL,Yr) 151 let Yb: *i64 = sys_mmap(T*dm*8) as *i64; stack_fwd(tape,vals,st,X,qq,qk,qv,qo,qg,qu,qd,T,dm,ffn,scale,NL,Yb) 152 let Yp: *i64 = sys_mmap(T*dm*8) as *i64; stack_fwd(tape,vals,st,X,pq,pk,pv,po,pg,pu,pd,T,dm,ffn,scale,NL,Yp) 153 let err_bf: i64 = l1err(Yb, Yr, T*dm) 154 let err_pt: i64 = l1err(Yp, Yr, T*dm) 155 156 // depth-1 (single block) per-tensor error, to show compounding with depth 157 let Yr1: *i64 = sys_mmap(T*dm*8) as *i64; stack_fwd(tape,vals,st,X,Wq,Wk,Wv,Wo,Wg,Wu,Wd,T,dm,ffn,scale,1,Yr1) 158 let Yp1: *i64 = sys_mmap(T*dm*8) as *i64; stack_fwd(tape,vals,st,X,pq,pk,pv,po,pg,pu,pd,T,dm,ffn,scale,1,Yp1) 159 let err_pt1: i64 = l1err(Yp1, Yr1, T*dm) 160 161 bb_puts(" "); bb_pn(NL); bb_puts("-block stack error: BLOCK-FLOAT="); bb_pn(err_bf); bb_puts(" PER-TENSOR="); bb_pn(err_pt) 162 bb_puts(" (per-tensor 1-block="); bb_pn(err_pt1); bb_puts(" -> "); bb_pn(NL); bb_puts("-block="); bb_pn(err_pt); bb_puts(")\n") 163 164 let Yb2: *i64 = sys_mmap(T*dm*8) as *i64; stack_fwd(tape,vals,st,X,qq,qk,qv,qo,qg,qu,qd,T,dm,ffn,scale,NL,Yb2) 165 var det: i64 = 1; var d: i64 = 0 166 while d < T*dm { if Yb2[d] != Yb[d] { det = 0 } d = d + 1 } 167 var nontriv: i64 = 0; var z: i64 = 0 168 while z < T*dm { if Yr[z] != 0 { nontriv = 1 } z = z + 1 } 169 170 var pass: i64 = 0; var total: i64 = 0 171 total = total + 1; pass = pass + bb_chk("T1 the stacked transformer RUNS with block-float weights" as *u8, nontriv) 172 var t2: i64 = 0; if err_bf < err_pt { t2 = 1 } 173 total = total + 1; pass = pass + bb_chk("T2 EXCEED: block-float stack error < per-tensor over the stack" as *u8, t2) 174 total = total + 1; pass = pass + bb_chk("T3 DETERMINISTIC: block-float stack twice == bit-identical" as *u8, det) 175 var t4: i64 = 0; if err_pt > err_pt1 { t4 = 1 } 176 total = total + 1; pass = pass + bb_chk("T4 DEPTH: per-tensor error compounds with depth (block-float stays stable)" as *u8, t4) 177 178 bb_puts("NX-BLOCKFLOAT-STACK-GATE "); bb_pn(pass); bb_puts(" / "); bb_pn(total) 179 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check 180 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled 181 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify. 182 let ctr__dry: *i64 = gv_ctr() 183 ctr__dry[0] = pass 184 ctr__dry[1] = total 185 let rc__dry: i64 = gv_verdict("BLOCKFLOAT-STACK-GATE" as *u8, ctr__dry, "block-float composes through depth: stable where per-tensor degrades -- toward a full model)" as *u8) 186 sys_exit(rc__dry) 187 return rc__dry 188}