code wiki / _hdl_build / nx_nofloat_multihead_gate.nx

nx_nofloat_multihead_gate.nx source

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