code wiki / _hdl_build / nx_nofloat_quant_gate.nx
nx_nofloat_quant_gate.nx source
↩ module page · 163 lines · 10643 B
1// nx_nofloat_quant_gate.nx -- CAP-NF-QUANT: post-training QUANTIZATION for efficient sovereign inference,
2// grounded in the BitNet research (nfs_bitnet158.raw / nfs_intquant.raw: ternary {-1,0,1} and integer-only
3// inference match full precision). Train the near-optimal 3-category grammar LM (Q16), then quantize its
4// weights to INT8 (per-tensor symmetric, scale=maxabs/127) and TERNARY {-1,0,1} (scale=mean|W|), and measure
5// how much held-out quality (CE) survives.
6// T1 INT8 preserves quality: int8 held-out CE ~= full-precision CE (within ~0.15 nats).
7// T2 INT8 stays near-OPTIMAL: int8 CE << uniform ln(8)=2079 (efficient 8-bit inference works).
8// T3 TERNARY carries the signal: ternary CE < uniform (1.58-bit weights still model the grammar; honest:
9// ternary loses more than int8 on a tiny model -- the BitNet "matches at scale" claim is a SCALE result).
10// Pure integer Q16. Sovereign: nx_nofloat_autograd + nx_syscalls. expect_exit: 0
11import "nx_nofloat_autograd.nx"
12import "nx_syscalls.nx"
13import "nx_gate_emit_lib.nx"
14const Q16: i64 = 65536
15const UNIFORM_MNAT: i64 = 2079 // ln(8)
16const FLOOR_MNAT: i64 = 964 // avg(ln2,ln3,ln3)
17
18
19func dini(a: *i64, n: i64, sd: i64) -> i64 { var i: i64=0; while i<n { a[i]=(((i*7+sd*13+1)%11)-5)*13107; i=i+1 } return 0 }
20func lcg(st: *i64) -> i64 { st[0]=(st[0]*1103515245 + 12345) & 2147483647; return (st[0] >> 15) }
21func iabs(x: i64) -> i64 { if x<0 { return 0-x } return x }
22func qround(src: i64, scale: i64) -> i64 { if scale<=0 { return 0 } if src>=0 { return (src + scale/2)/scale } return (src - scale/2)/scale }
23// INT8 per-tensor symmetric quantize+dequantize: src(Q16) -> 8-bit levels [-127,127] -> dequantized(Q16)
24func quant_int8(src: *i64, dst: *i64, n: i64) -> i64 {
25 var mx: i64=0; var i: i64=0; while i<n { let a: i64=iabs(src[i]); if a>mx { mx=a } i=i+1 }
26 if mx==0 { i=0; while i<n { dst[i]=0; i=i+1 } return 0 }
27 var scale: i64=mx/127; if scale==0 { scale=1 }
28 i=0; while i<n { var q: i64=qround(src[i],scale); if q>127 { q=127 } if q<0-127 { q=0-127 } dst[i]=q*scale; i=i+1 }
29 return 0
30}
31// TERNARY {-1,0,1} quantize+dequantize (BitNet b1.58 style: scale = mean|W|)
32func quant_ternary(src: *i64, dst: *i64, n: i64) -> i64 {
33 var s: i64=0; var i: i64=0; while i<n { s=s+iabs(src[i]); i=i+1 }
34 var scale: i64=s/n; if scale==0 { scale=1 }
35 i=0; while i<n { var q: i64=qround(src[i],scale); if q>1 { q=1 } if q<0-1 { q=0-1 } dst[i]=q*scale; i=i+1 }
36 return 0
37}
38func make_stream(S: *i64, tgt: *i64, P: i64, st: *i64) -> i64 {
39 var i: i64=0
40 while i<P { let c: i64=i%3; if c==0 { S[i]=lcg(st)%2 } if c==1 { S[i]=2+lcg(st)%3 } if c==2 { S[i]=5+lcg(st)%3 } i=i+1 }
41 i=0; while i<P-1 { tgt[i]=S[i+1]; i=i+1 } tgt[P-1]=S[0]
42 return 0
43}
44func clm_fwd(tape: *i64, vals: *i64, st: *i64, W: *i64, ids: *i64, tgt: *i64, T: i64, dm: i64, V: i64, scale: i64, leaves: *i64) -> i64 {
45 let E: *i64=W[0] as *i64; let Wq: *i64=W[1] as *i64; let Wk: *i64=W[2] as *i64; let Wv: *i64=W[3] as *i64; let Wo: *i64=W[4] as *i64; let Wlm: *i64=W[5] as *i64
46 st[0]=0; st[1]=0
47 let nE: i64=nfa_leaf(tape,vals,st,V,dm,E,0)
48 let nWq: i64=nfa_leaf(tape,vals,st,dm,dm,Wq,0)
49 let nWk: i64=nfa_leaf(tape,vals,st,dm,dm,Wk,0)
50 let nWv: i64=nfa_leaf(tape,vals,st,dm,dm,Wv,0)
51 let nWo: i64=nfa_leaf(tape,vals,st,dm,dm,Wo,0)
52 let nWlm: i64=nfa_leaf(tape,vals,st,dm,V,Wlm,0)
53 let nX: i64=nfa_embed(tape,vals,st,nE,ids,T)
54 let nXn: i64=nfa_rmsnorm_rows(tape,vals,st,nX)
55 let nQ: i64=nfa_matmul(tape,vals,st,nXn,nWq)
56 let nK: i64=nfa_matmul(tape,vals,st,nXn,nWk)
57 let nV: i64=nfa_matmul(tape,vals,st,nXn,nWv)
58 let nQr: i64=nfa_rope(tape,vals,st,nQ)
59 let nKr: i64=nfa_rope(tape,vals,st,nK)
60 let nS: i64=nfa_matmul_nt(tape,vals,st,nQr,nKr)
61 let nSs: i64=nfa_cmul(tape,vals,st,nS,scale)
62 let nA: i64=nfa_softmax_rows(tape,vals,st,nSs,1)
63 let nO: i64=nfa_matmul(tape,vals,st,nA,nV)
64 let nOp: i64=nfa_matmul(tape,vals,st,nO,nWo)
65 let nH: i64=nfa_vadd(tape,vals,st,nX,nOp)
66 let nHn: i64=nfa_rmsnorm_rows(tape,vals,st,nH)
67 let nLg: i64=nfa_matmul(tape,vals,st,nHn,nWlm)
68 let nLoss: i64=nfa_softce_rows(tape,vals,st,nLg,tgt)
69 leaves[0]=nE; leaves[1]=nWq; leaves[2]=nWk; leaves[3]=nWv; leaves[4]=nWo; leaves[5]=nWlm; leaves[6]=nLg
70 return nLoss
71}
72func step_all(tape: *i64, grads: *i64, W: *i64, WN: *i64, leaves: *i64, lr: i64, clip: i64, gb: *i64) -> i64 {
73 var i: i64=0
74 while i<6 { let ar: *i64=W[i] as *i64; let cn: i64=WN[i]; let nd: i64=leaves[i]; var c: i64=0; while c<cn { var g: i64=nfa_grad(tape,grads,nd,c); if g>clip{g=clip} if g<0-clip{g=0-clip} gb[c]=g; c=c+1 } nfa_sgd(ar,gb,cn,lr); i=i+1 }
75 return 0
76}
77func do_train(tape: *i64, vals: *i64, grads: *i64, st: *i64, W: *i64, WN: *i64, S: *i64, tgt: *i64, P: i64, dm: i64, V: i64, scale: i64, leaves: *i64, gb: *i64, steps: i64, sdat: *i64) -> i64 {
78 var ep: i64=0
79 while ep < steps { make_stream(S,tgt,P,sdat); let nl: i64=clm_fwd(tape,vals,st,W,S,tgt,P-1,dm,V,scale,leaves); nfa_backward(tape,vals,grads,st[0],nl); step_all(tape,grads,W,WN,leaves,6554,262144,gb); ep=ep+1 }
80 return 0
81}
82// ternary QAT (BitNet b1.58 recipe): forward uses ternary-quantized weights (Wt), gradients update the
83// full-precision SHADOW weights W2 (straight-through estimator: quantization treated as identity in backward).
84// So W2 learns to be ROBUST to ternary quantization -- which naive post-training ternary can't be.
85func do_train_qat(tape: *i64, vals: *i64, grads: *i64, st: *i64, W2: *i64, Wt: *i64, WN: *i64, S: *i64, tgt: *i64, P: i64, dm: i64, V: i64, scale: i64, leaves: *i64, gb: *i64, steps: i64, sdat: *i64) -> i64 {
86 var ep: i64=0
87 while ep < steps {
88 make_stream(S,tgt,P,sdat)
89 var i: i64=0; while i<6 { quant_ternary(W2[i] as *i64, Wt[i] as *i64, WN[i]); i=i+1 } // quantize shadow -> ternary
90 let nl: i64=clm_fwd(tape,vals,st,Wt,S,tgt,P-1,dm,V,scale,leaves) // forward with ternary weights
91 nfa_backward(tape,vals,grads,st[0],nl)
92 step_all(tape,grads,W2,WN,leaves,6554,262144,gb) // STE: grads(Wt) -> update shadow W2
93 ep=ep+1
94 }
95 return 0
96}
97func eval_ce(tape: *i64, vals: *i64, st: *i64, W: *i64, S: *i64, tgt: *i64, P: i64, dm: i64, V: i64, scale: i64, leaves: *i64, N: i64, sdat: *i64) -> i64 {
98 var acc: i64=0; var e: i64=0
99 while e<N { make_stream(S,tgt,P,sdat); let nl: i64=clm_fwd(tape,vals,st,W,S,tgt,P-1,dm,V,scale,leaves); acc=acc+nfa_val(tape,vals,nl,0); e=e+1 }
100 let mq: i64=acc/N
101 return (mq*1000)/Q16
102}
103
104func main() -> i64 {
105 g_puts("nx_nofloat_quant gate (post-training INT8 + TERNARY quantization for efficient sovereign inference)\n" as *u8)
106 let V: i64=8; let P: i64=12; let dm: i64=24; let scale: i64=13377
107 let tape: *i64=sys_mmap(512*7*8) as *i64
108 let vals: *i64=sys_mmap(65536*8) as *i64
109 let grads: *i64=sys_mmap(65536*8) as *i64
110 let st: *i64=sys_mmap(2*8) as *i64
111 let nW: i64=6
112 let W: *i64=sys_mmap(nW*8) as *i64; let WN: *i64=sys_mmap(nW*8) as *i64
113 WN[0]=V*dm; WN[1]=dm*dm; WN[2]=dm*dm; WN[3]=dm*dm; WN[4]=dm*dm; WN[5]=dm*V
114 var wi: i64=0; while wi<nW { let a: *i64=sys_mmap(WN[wi]*8) as *i64; dini(a,WN[wi],wi+1); W[wi]=a as i64; wi=wi+1 }
115 // a parallel set of dequantized-weight arrays + its pointer table
116 let Wq: *i64=sys_mmap(nW*8) as *i64
117 wi=0; while wi<nW { Wq[wi]=(sys_mmap(WN[wi]*8) as *i64) as i64; wi=wi+1 }
118 // QAT shadow weights (fresh init, trained ternary-aware)
119 let W2: *i64=sys_mmap(nW*8) as *i64
120 wi=0; while wi<nW { let a: *i64=sys_mmap(WN[wi]*8) as *i64; dini(a,WN[wi],wi+1); W2[wi]=a as i64; wi=wi+1 }
121 let leaves: *i64=sys_mmap(8*8) as *i64; let gbuf: *i64=sys_mmap(4096*8) as *i64
122 let S: *i64=sys_mmap(P*8) as *i64; let tgt: *i64=sys_mmap(P*8) as *i64; let sdat: *i64=sys_mmap(8) as *i64
123
124 sdat[0]=12345
125 do_train(tape,vals,grads,st,W,WN,S,tgt,P,dm,V,scale,leaves,gbuf,30000,sdat)
126
127 sdat[0]=70707070
128 let full_ce: i64=eval_ce(tape,vals,st,W,S,tgt,P,dm,V,scale,leaves,200,sdat)
129 // quantize all weights to int8 -> Wq, eval
130 wi=0; while wi<nW { quant_int8(W[wi] as *i64, Wq[wi] as *i64, WN[wi]); wi=wi+1 }
131 sdat[0]=70707070
132 let int8_ce: i64=eval_ce(tape,vals,st,Wq,S,tgt,P,dm,V,scale,leaves,200,sdat)
133 // naive post-training ternary (expected to FAIL -- ternary needs quantization-AWARE training)
134 wi=0; while wi<nW { quant_ternary(W[wi] as *i64, Wq[wi] as *i64, WN[wi]); wi=wi+1 }
135 sdat[0]=70707070
136 let tern_ptq_ce: i64=eval_ce(tape,vals,st,Wq,S,tgt,P,dm,V,scale,leaves,200,sdat)
137
138 // ternary QAT: train the shadow W2 ternary-aware (STE), then eval its ternary-quantized weights
139 sdat[0]=12345
140 do_train_qat(tape,vals,grads,st,W2,Wq,WN,S,tgt,P,dm,V,scale,leaves,gbuf,30000,sdat)
141 wi=0; while wi<nW { quant_ternary(W2[wi] as *i64, Wq[wi] as *i64, WN[wi]); wi=wi+1 }
142 sdat[0]=70707070
143 let tern_qat_ce: i64=eval_ce(tape,vals,st,Wq,S,tgt,P,dm,V,scale,leaves,200,sdat)
144
145 g_puts(" [measure] held-out CE (milli-nats): full-Q16="); g_pn(full_ce); g_puts(" INT8-PTQ="); g_pn(int8_ce); g_puts(" TERNARY-PTQ="); g_pn(tern_ptq_ce); g_puts(" (naive, fails) TERNARY-QAT="); g_pn(tern_qat_ce); g_puts(" (uniform="); g_pn(UNIFORM_MNAT); g_puts(", floor="); g_pn(FLOOR_MNAT); g_puts(")\n")
146
147 var pass: i64=0; var total: i64=0
148 var t1: i64=0; if int8_ce <= full_ce+150 { t1=1 }
149 pass=pass+g_check("T1: INT8 post-training quant preserves quality (int8 CE ~= full-precision CE)" as *u8, t1); total=total+1
150 var t2: i64=0; if int8_ce*10 <= UNIFORM_MNAT*7 { t2=1 }
151 pass=pass+g_check("T2: INT8 stays near-OPTIMAL (<< uniform = efficient 8-bit sovereign inference works)" as *u8, t2); total=total+1
152 var t3: i64=0; if tern_qat_ce < UNIFORM_MNAT { if tern_qat_ce*2 < tern_ptq_ce { t3=1 } }
153 pass=pass+g_check("T3: ternary QAT >> naive PTQ (BitNet insight: ternary needs quant-AWARE training; QAT CE < uniform AND << PTQ)" as *u8, t3); total=total+1
154
155 var okall: i64=0; if pass==total { okall=1 }
156 if okall==1 {
157 let logf: i64=sys_openat_append("knowledge/status/nofloat_quant.log" as *u8, 420)
158 if logf>=0 { let x0: i64=sys_write(logf,"NOFLOATQUANT int8+ternary post-training quantization measured\n" as *u8,62); sys_close(logf) }
159 }
160 g_puts("---- quant gate: passed "); g_pn(pass); g_puts(" / "); g_pn(total); g_puts(" ----\n")
161 if okall==1 { g_puts("verdict=GREEN (post-training INT8 inference preserves quality; ternary carries signal -- efficient sovereign inference)\n" as *u8); sys_exit(0); return 0 }
162 g_puts("verdict=RED\n" as *u8); sys_exit(1); return 1
163}