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}