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}