code wiki / _hdl_build / nx_nofloat_multihead_gate.nx

nx_nofloat_multihead_gate.nx source

↩ module page · 260 lines · 15939 B

1// nx_nofloat_multihead_gate.nx -- HARD-EVIDENCE gate for MULTI-HEAD attention in pure integer Q16. The defining 2// transformer feature: split Q/K/V into H heads (slice_cols), run the verified attention core per head (RoPE + 3// scaled causal softmax + A.V), concat the heads (concat_cols), then the output projection. Backprops end-to-end. 4// 5// A1 slice_cols gradcheck : extract cols backward (scatter) == finite differences. 6// A2 concat_cols gradcheck : concat backward (split) == finite differences. 7// A3 multi-head gradcheck wrt Wq : the gradient through BOTH heads (slice+RoPE+softmax+A.V+concat+proj) == FD. 8// D neg-control teeth ; C bit-exact. 9// B multi-head LEARNS : with attention fixed, the output projection Wo is trained (AdamW) to a realizable 10// target; loss collapses -> the multi-head output path trains. 11// 12// Evidence -> knowledge/status/nofloat_multihead.log. Sovereign: nx_nofloat_autograd + nx_syscalls. expect_exit: 0 13import "nx_nofloat_autograd.nx" 14import "nx_syscalls.nx" 15import "nx_gate_emit_lib.nx" 16import "nx_gate_verdict.nx" 17 18const MLOG: *u8 = "knowledge/status/nofloat_multihead.log" 19const Q16: i64 = 65536 20 21 22func g_abs(v: i64) -> i64 { if v < 0 { return 0 - v } return v } 23func m_ws(fd: i64, s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} sys_write(fd,s,n); return 0 } 24func m_wn(fd: i64, v: i64) -> i64 { let b: *u8=sys_mmap(28); var m: i64=v; if m<0{sys_write(fd,"-" as *u8,1);m=0-m} let t: *u8=sys_mmap(28); var k: i64=0; if m==0{t[0]=48;k=1} while m>0{t[k]=(48+(m%10)) as u8;m=m/10;k=k+1} var i: i64=0; while i<k{b[i]=t[k-1-i];i=i+1} sys_write(fd,b,k); return 0 } 25// init scale 13107 (~[-1,1]): big enough that attention logits spread (non-uniform softmax) -> real Wq gradients 26// (RMSNorm makes the INPUT unit-scale regardless, so the logit spread comes from the Wq/Wk magnitude, not X). 27func lm_init(arr: *i64, n: i64, seed: i64) -> i64 { var i: i64=0; while i<n { arr[i] = (((i*7 + seed*13 + 1) % 11) - 5) * 13107; i=i+1 } return 0 } 28 29// ---- multi-head attention forward; W=[Wq,Wk,Wv,Wo]; leaves[0..3]; returns out node ---- 30func mha_fwd(tape: *i64, vals: *i64, st: *i64, W: *i64, X: *i64, T: i64, dm: i64, hd: i64, H: i64, scale: i64, leaves: *i64) -> i64 { 31 let Wq: *i64 = W[0] as *i64; let Wk: *i64 = W[1] as *i64; let Wv: *i64 = W[2] as *i64; let Wo: *i64 = W[3] as *i64 32 st[0]=0; st[1]=0 33 let nX: i64 = nfa_leaf(tape,vals,st,T,dm,X,0) 34 let nWq: i64 = nfa_leaf(tape,vals,st,dm,dm,Wq,0) 35 let nWk: i64 = nfa_leaf(tape,vals,st,dm,dm,Wk,0) 36 let nWv: i64 = nfa_leaf(tape,vals,st,dm,dm,Wv,0) 37 let nWo: i64 = nfa_leaf(tape,vals,st,dm,dm,Wo,0) 38 let nXn: i64 = nfa_rmsnorm_rows(tape,vals,st,nX) 39 let nQ: i64 = nfa_matmul(tape,vals,st,nXn,nWq) 40 let nK: i64 = nfa_matmul(tape,vals,st,nXn,nWk) 41 let nV: i64 = nfa_matmul(tape,vals,st,nXn,nWv) 42 var Oacc: i64 = 0 - 1 43 var hh: i64 = 0 44 while hh < H { 45 let nQh: i64 = nfa_slice_cols(tape,vals,st,nQ,hh*hd,hd) 46 let nKh: i64 = nfa_slice_cols(tape,vals,st,nK,hh*hd,hd) 47 let nVh: i64 = nfa_slice_cols(tape,vals,st,nV,hh*hd,hd) 48 let nQr: i64 = nfa_rope(tape,vals,st,nQh) 49 let nKr: i64 = nfa_rope(tape,vals,st,nKh) 50 let nS: i64 = nfa_matmul_nt(tape,vals,st,nQr,nKr) 51 let nSs: i64 = nfa_cmul(tape,vals,st,nS,scale) 52 let nA: i64 = nfa_softmax_rows(tape,vals,st,nSs,1) 53 let nOh: i64 = nfa_matmul(tape,vals,st,nA,nVh) 54 if Oacc < 0 { Oacc = nOh } else { Oacc = nfa_concat_cols(tape,vals,st,Oacc,nOh) } 55 hh = hh + 1 56 } 57 let nOp: i64 = nfa_matmul(tape,vals,st,Oacc,nWo) 58 let nOut: i64 = nfa_vadd(tape,vals,st,nX,nOp) 59 leaves[0]=nWq; leaves[1]=nWk; leaves[2]=nWv; leaves[3]=nWo 60 return nOut 61} 62func mha_loss(tape: *i64, vals: *i64, st: *i64, W: *i64, X: *i64, Tg: *i64, T: i64, dm: i64, hd: i64, H: i64, scale: i64, leaves: *i64) -> i64 { 63 let nOut: i64 = mha_fwd(tape,vals,st,W,X,T,dm,hd,H,scale,leaves) 64 let nt: i64 = nfa_leaf(tape,vals,st,T,dm,Tg,0) 65 return nfa_mse(tape,vals,st,nOut,nt) 66} 67func mha_lossval(tape: *i64, vals: *i64, st: *i64, W: *i64, X: *i64, Tg: *i64, T: i64, dm: i64, hd: i64, H: i64, scale: i64) -> i64 { 68 let lv: *i64 = sys_mmap(4*8) as *i64 69 let loss: i64 = mha_loss(tape,vals,st,W,X,Tg,T,dm,hd,H,scale,lv) 70 return nfa_val(tape,vals,loss,0) 71} 72func mha_out(tape: *i64, vals: *i64, st: *i64, W: *i64, X: *i64, T: i64, dm: i64, hd: i64, H: i64, scale: i64, outv: *i64) -> i64 { 73 let lv: *i64 = sys_mmap(4*8) as *i64 74 let nOut: i64 = mha_fwd(tape,vals,st,W,X,T,dm,hd,H,scale,lv) 75 var i: i64=0 76 while i<T*dm { outv[i]=nfa_val(tape,vals,nOut,i); i=i+1 } 77 return 0 78} 79 80func main() -> i64 { 81 g_puts("nx_nofloat_multihead gate (H-head causal attention backprops + trains, PURE INTEGER Q16)\n" as *u8) 82 var pass: i64=0; var total: i64=0 83 let tape: *i64 = sys_mmap(512*7*8) as *i64 84 let vals: *i64 = sys_mmap(8192*8) as *i64 85 let grads: *i64 = sys_mmap(8192*8) as *i64 86 let st: *i64 = sys_mmap(2*8) as *i64 87 let h: i64 = 512; let floor_q: i64 = 4096 88 89 // ---- A1: slice_cols gradcheck (X[2,4], slice c0=1,w=2) ---- 90 let sX: *i64 = sys_mmap(8*8) as *i64; lm_init(sX,8,3) 91 let sT: *i64 = sys_mmap(4*8) as *i64; sT[0]=6554; sT[1]=0-13107; sT[2]=19661; sT[3]=3277 92 var s_ok: i64=1; var s_worst: i64=0 93 st[0]=0; st[1]=0 94 let snX: i64 = nfa_leaf(tape,vals,st,2,4,sX,0) 95 let snS: i64 = nfa_slice_cols(tape,vals,st,snX,1,2) 96 let snT: i64 = nfa_leaf(tape,vals,st,2,2,sT,0) 97 let sloss: i64 = nfa_mse(tape,vals,st,snS,snT) 98 nfa_backward(tape,vals,grads,st[0],sloss) 99 let sana: *i64 = sys_mmap(8*8) as *i64 100 var sc: i64=0 101 while sc<8 { sana[sc]=nfa_grad(tape,grads,snX,sc); sc=sc+1 } 102 var si: i64=0 103 while si<8 { 104 let old: i64=sX[si] 105 sX[si]=old+h 106 st[0]=0; st[1]=0; let a1: i64=nfa_leaf(tape,vals,st,2,4,sX,0); let a2: i64=nfa_slice_cols(tape,vals,st,a1,1,2); let a3: i64=nfa_leaf(tape,vals,st,2,2,sT,0); let lp: i64=nfa_val(tape,vals,nfa_mse(tape,vals,st,a2,a3),0) 107 sX[si]=old-h 108 st[0]=0; st[1]=0; let b1: i64=nfa_leaf(tape,vals,st,2,4,sX,0); let b2: i64=nfa_slice_cols(tape,vals,st,b1,1,2); let b3: i64=nfa_leaf(tape,vals,st,2,2,sT,0); let lm2: i64=nfa_val(tape,vals,nfa_mse(tape,vals,st,b2,b3),0) 109 sX[si]=old 110 let fd: i64=((lp-lm2)*Q16)/(2*h); let num: i64=g_abs(fd-sana[si]); var den: i64=g_abs(sana[si]); if den<floor_q{den=floor_q} 111 if num >= ((4096*den)>>16) { s_ok=0 } 112 let rel: i64=(num*1000)/den; if rel>s_worst{s_worst=rel} 113 si=si+1 114 } 115 g_puts(" [measure] slice_cols worst rel grad err = " as *u8); g_pn(s_worst); g_puts(" /1000 (tol=62)\n" as *u8) 116 pass=pass+g_check("A1: slice_cols gradcheck -- column extract backward (scatter) == finite differences" as *u8, s_ok); total=total+1 117 118 // ---- A2: concat_cols gradcheck (a[2,2], b[2,3]) wrt a ---- 119 let ca: *i64 = sys_mmap(4*8) as *i64; lm_init(ca,4,4) 120 let cb: *i64 = sys_mmap(6*8) as *i64; lm_init(cb,6,5) 121 let cT: *i64 = sys_mmap(10*8) as *i64; var ct: i64=0; while ct<10 { cT[ct]=(ct-5)*4096; ct=ct+1 } 122 var co_ok: i64=1; var co_worst: i64=0 123 st[0]=0; st[1]=0 124 let cna: i64 = nfa_leaf(tape,vals,st,2,2,ca,0) 125 let cnb: i64 = nfa_leaf(tape,vals,st,2,3,cb,0) 126 let cnc: i64 = nfa_concat_cols(tape,vals,st,cna,cnb) 127 let cnt: i64 = nfa_leaf(tape,vals,st,2,5,cT,0) 128 let closs: i64 = nfa_mse(tape,vals,st,cnc,cnt) 129 nfa_backward(tape,vals,grads,st[0],closs) 130 let cana: *i64 = sys_mmap(4*8) as *i64 131 var cc: i64=0 132 while cc<4 { cana[cc]=nfa_grad(tape,grads,cna,cc); cc=cc+1 } 133 var ci: i64=0 134 while ci<4 { 135 let old: i64=ca[ci] 136 ca[ci]=old+h 137 st[0]=0; st[1]=0; let a1: i64=nfa_leaf(tape,vals,st,2,2,ca,0); let a2: i64=nfa_leaf(tape,vals,st,2,3,cb,0); let a3: i64=nfa_concat_cols(tape,vals,st,a1,a2); let a4: i64=nfa_leaf(tape,vals,st,2,5,cT,0); let lp: i64=nfa_val(tape,vals,nfa_mse(tape,vals,st,a3,a4),0) 138 ca[ci]=old-h 139 st[0]=0; st[1]=0; let b1: i64=nfa_leaf(tape,vals,st,2,2,ca,0); let b2: i64=nfa_leaf(tape,vals,st,2,3,cb,0); let b3: i64=nfa_concat_cols(tape,vals,st,b1,b2); let b4: i64=nfa_leaf(tape,vals,st,2,5,cT,0); let lm2: i64=nfa_val(tape,vals,nfa_mse(tape,vals,st,b3,b4),0) 140 ca[ci]=old 141 let fd: i64=((lp-lm2)*Q16)/(2*h); let num: i64=g_abs(fd-cana[ci]); var den: i64=g_abs(cana[ci]); if den<floor_q{den=floor_q} 142 if num >= ((4096*den)>>16) { co_ok=0 } 143 let rel: i64=(num*1000)/den; if rel>co_worst{co_worst=rel} 144 ci=ci+1 145 } 146 g_puts(" [measure] concat_cols worst rel grad err = " as *u8); g_pn(co_worst); g_puts(" /1000 (tol=62)\n" as *u8) 147 pass=pass+g_check("A2: concat_cols gradcheck -- concat backward (split) == finite differences" as *u8, co_ok); total=total+1 148 149 // ---- MHA dims + weights ---- 150 let T: i64=3; let dm: i64=4; let hd: i64=2; let H: i64=2; let scale: i64=46341 // 1/sqrt(2) 151 let X: *i64 = sys_mmap(T*dm*8) as *i64; lm_init(X,T*dm,1) 152 let Wq: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wq,dm*dm,2) 153 let Wk: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wk,dm*dm,3) 154 let Wv: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wv,dm*dm,4) 155 let Wo: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wo,dm*dm,5) 156 let W: *i64 = sys_mmap(4*8) as *i64; W[0]=Wq as i64; W[1]=Wk as i64; W[2]=Wv as i64; W[3]=Wo as i64 157 let Tg: *i64 = sys_mmap(T*dm*8) as *i64; var tgi: i64=0; while tgi<T*dm { Tg[tgi]=((tgi%7)-3)*4096; tgi=tgi+1 } 158 let leaves: *i64 = sys_mmap(4*8) as *i64 159 160 // ---- A3: multi-head gradcheck wrt Wq (through BOTH heads) ---- 161 let l3: i64 = mha_loss(tape,vals,st,W,X,Tg,T,dm,hd,H,scale,leaves) 162 nfa_backward(tape,vals,grads,st[0],l3) 163 let nWq: i64 = leaves[0] 164 let a3ana: *i64 = sys_mmap(16*8) as *i64 165 var ac: i64=0 166 while ac<dm*dm { a3ana[ac]=nfa_grad(tape,grads,nWq,ac); ac=ac+1 } 167 var a3_ok: i64=1; var a3_worst: i64=0 168 var qi: i64=0 169 while qi<dm*dm { 170 let old: i64=Wq[qi]; Wq[qi]=old+h; let lp: i64=mha_lossval(tape,vals,st,W,X,Tg,T,dm,hd,H,scale); Wq[qi]=old-h; let lm2: i64=mha_lossval(tape,vals,st,W,X,Tg,T,dm,hd,H,scale); Wq[qi]=old 171 let fd: i64=((lp-lm2)*Q16)/(2*h); let num: i64=g_abs(fd-a3ana[qi]); var den: i64=g_abs(a3ana[qi]); if den<floor_q{den=floor_q} 172 if num >= ((16384*den)>>16) { a3_ok=0 } // tol 1/4 (multi-head chain through 2 softmaxes) 173 let rel: i64=(num*1000)/den; if rel>a3_worst{a3_worst=rel} 174 qi=qi+1 175 } 176 g_puts(" [measure] multi-head dL/dWq worst rel grad err = " as *u8); g_pn(a3_worst); g_puts(" /1000 (tol=250)\n" as *u8) 177 pass=pass+g_check("A3: multi-head gradcheck wrt Wq through BOTH heads (slice+RoPE+softmax+A.V+concat+proj)" as *u8, a3_ok); total=total+1 178 179 // ---- D: neg-control teeth (use the LARGEST-magnitude grad component so the test isn't vacuous on a ~0 grad) ---- 180 var imax: i64=0; var vmax: i64=g_abs(a3ana[0]); var ii: i64=1 181 while ii<dm*dm { if g_abs(a3ana[ii])>vmax { vmax=g_abs(a3ana[ii]); imax=ii } ii=ii+1 } 182 let dana: i64 = a3ana[imax] 183 let o0: i64=Wq[imax]; Wq[imax]=o0+h; let lpd: i64=mha_lossval(tape,vals,st,W,X,Tg,T,dm,hd,H,scale); Wq[imax]=o0-h; let lmd: i64=mha_lossval(tape,vals,st,W,X,Tg,T,dm,hd,H,scale); Wq[imax]=o0 184 let dfd: i64=((lpd-lmd)*Q16)/(2*h); let dbad: i64=0-dana 185 // scale-free teeth: the FD must be unambiguously CLOSER to the true analytic grad than to the negated one 186 // (works even when grads are tiny, where an absolute floor would swamp the signal). vmax guards vs pure noise. 187 let dgood: i64 = g_abs(dfd - dana); let dneg: i64 = g_abs(dfd - dbad) 188 var caught: i64=0; if vmax > 64 { if dneg > dgood*4 { caught=1 } } 189 pass=pass+g_check("D: neg-control -- FD is >4x closer to the true grad than to the negated one (teeth)" as *u8, caught); total=total+1 190 191 // ---- B: multi-head output trains (fix Wq/Wk/Wv -> attention+values fixed -> Op=Oacc.Wo convex; AdamW) ---- 192 let Wot: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wot,dm*dm,9) 193 let Wt2: *i64 = sys_mmap(4*8) as *i64; Wt2[0]=Wq as i64; Wt2[1]=Wk as i64; Wt2[2]=Wv as i64; Wt2[3]=Wot as i64 194 let tgt: *i64 = sys_mmap(T*dm*8) as *i64 195 mha_out(tape,vals,st,Wt2,X,T,dm,hd,H,scale,tgt) // realizable target = mha out with Wo* 196 let Wop: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wop,dm*dm,2) // student Wo (different init) 197 let Wb: *i64 = sys_mmap(4*8) as *i64; Wb[0]=Wq as i64; Wb[1]=Wk as i64; Wb[2]=Wv as i64; Wb[3]=Wop as i64 198 let gWo: *i64 = sys_mmap(16*8) as *i64 199 let mm: *i64 = sys_mmap(16*8) as *i64; let vv: *i64 = sys_mmap(16*8) as *i64 200 var zi: i64=0; while zi<dm*dm { mm[zi]=0; vv[zi]=0; zi=zi+1 } 201 let lvb: *i64 = sys_mmap(4*8) as *i64 202 var lf: i64=0; var ll: i64=0 203 var ep: i64=0 204 while ep < 2000 { 205 let lb: i64 = mha_loss(tape,vals,st,Wb,X,tgt,T,dm,hd,H,scale,lvb) 206 nfa_backward(tape,vals,grads,st[0],lb) 207 if ep==0 { lf=nfa_val(tape,vals,lb,0) } 208 ll=nfa_val(tape,vals,lb,0) 209 let nWoL: i64 = lvb[3] 210 var z: i64=0 211 while z<dm*dm { gWo[z]=nfa_grad(tape,grads,nWoL,z); z=z+1 } 212 nfa_adamw(Wop, gWo, mm, vv, dm*dm, 3277, 58982, 65470, 66, 0, ep+1) 213 ep=ep+1 214 } 215 g_puts(" [measure] multi-head Wo-train loss: start=" as *u8); g_pn(lf); g_puts(" end=" as *u8); g_pn(ll); g_puts("\n" as *u8) 216 var learns: i64=1 217 if ll*5 > lf { learns=0 } // >= 80% loss reduction 218 if lf<=0 { learns=0 } 219 pass=pass+g_check("B: multi-head output LEARNS -- Wo trained with AdamW to a realizable target, loss collapses" as *u8, learns); total=total+1 220 221 // ---- C: bit-exact ---- 222 let Wop2: *i64 = sys_mmap(dm*dm*8) as *i64; lm_init(Wop2,dm*dm,2) 223 let Wc: *i64 = sys_mmap(4*8) as *i64; Wc[0]=Wq as i64; Wc[1]=Wk as i64; Wc[2]=Wv as i64; Wc[3]=Wop2 as i64 224 let gWo2: *i64 = sys_mmap(16*8) as *i64 225 let mm2: *i64 = sys_mmap(16*8) as *i64; let vv2: *i64 = sys_mmap(16*8) as *i64 226 zi=0; while zi<dm*dm { mm2[zi]=0; vv2[zi]=0; zi=zi+1 } 227 let lvc: *i64 = sys_mmap(4*8) as *i64 228 var ep2: i64=0 229 while ep2 < 2000 { 230 let lc: i64 = mha_loss(tape,vals,st,Wc,X,tgt,T,dm,hd,H,scale,lvc) 231 nfa_backward(tape,vals,grads,st[0],lc) 232 let nWoL: i64=lvc[3]; var z: i64=0 233 while z<dm*dm { gWo2[z]=nfa_grad(tape,grads,nWoL,z); z=z+1 } 234 nfa_adamw(Wop2, gWo2, mm2, vv2, dm*dm, 3277, 58982, 65470, 66, 0, ep2+1) 235 ep2=ep2+1 236 } 237 var bitexact: i64=1 238 var bz: i64=0 239 while bz<dm*dm { if Wop2[bz]!=Wop[bz] { bitexact=0 } bz=bz+1 } 240 pass=pass+g_check("C: bit-exact -- training the multi-head twice gives IDENTICAL integer Wo (determinism)" as *u8, bitexact); total=total+1 241 242 var okall: i64=0; if pass==total { okall=1 } 243 let logf: i64 = sys_openat_append(MLOG, 420) 244 if logf >= 0 { 245 m_ws(logf,"NOFLOATMULTIHEAD H=2 hd=2 A1_slice=" as *u8); m_wn(logf,s_ok); m_ws(logf," A2_concat=" as *u8); m_wn(logf,co_ok) 246 m_ws(logf," A3_mha_Wq=" as *u8); m_wn(logf,a3_ok); m_ws(logf," D=" as *u8); m_wn(logf,caught); m_ws(logf," B_learns=" as *u8); m_wn(logf,learns); m_ws(logf," C_bitexact=" as *u8); m_wn(logf,bitexact) 247 if okall==1 { m_ws(logf," verdict=GREEN\n" as *u8) } else { m_ws(logf," verdict=RED\n" as *u8) } 248 sys_close(logf) 249 } 250 g_puts("---- nofloat_multihead gate: passed " as *u8); g_pn(pass); g_puts(" / " as *u8); g_pn(total); g_puts(" ----\n" as *u8) 251 // MIGRATED onto nx_gate_verdict by nx_gate_dry_apply (D001, minimal form): every check 252 // row above is untouched, so the PASS/FAIL vector cannot change; only the hand-rolled 253 // verdict emission is replaced by the ONE shared base class. Proven by nx_gate_migrate verify. 254 let ctr__dry: *i64 = gv_ctr() 255 ctr__dry[0] = pass 256 ctr__dry[1] = total 257 let rc__dry: i64 = gv_verdict("NOFLOAT-MULTIHEAD-GATE" as *u8, ctr__dry, "multi-head causal attention backprops + trains end-to-end in pure integer Q16)" as *u8) 258 sys_exit(rc__dry) 259 return rc__dry 260}