code wiki / _hdl_build / nx_nofloat_attn_gate.nx

nx_nofloat_attn_gate.nx source

↩ module page · 244 lines · 13804 B

1// nx_nofloat_attn_gate.nx -- HARD-EVIDENCE gate for the ATTENTION-CORE backward (CAP-NF-ATTN-CORE): a single-head 2// CAUSAL self-attention block backprops end-to-end in PURE INTEGER Q16. This is the heart of CAP-NF-TRAIN-ATTN. 3// New ops proven: matmul (C=A.B), matmul_nt (S=Q.K^T), cmul (1/sqrt(d) scale), causal softmax_rows. 4// 5// A1 matmul gradcheck : C=A.B, loss=mse(C,t); tape grad dL/dA == central finite difference. 6// A2 matmul_nt gradcheck : S=A.B^T (the Q.K^T form); tape grad == finite difference. 7// A3 attention-core gradcheck: full path X->{Q,K,V}=X.W -> S=Q.K^T -> scale -> CAUSAL softmax rows -> O=A.V -> 8// mse; gradcheck dL/dWq, the gradient that flows THROUGH softmax + both matmuls (the real attention backward). 9// D neg-control teeth : a deliberately wrong matmul grad is rejected. 10// B attention LEARNS : with Wq,Wk fixed (attention pattern A fixed), the VALUE path O=A.(X.Wv) is linear -> 11// train Wv from zero to a realizable target; assert loss collapses + Wv converges (the value projection trains). 12// C bit-exact : train twice -> identical integer Wv (determinism is structural for integer). 13// 14// Evidence -> knowledge/status/nofloat_attn.log. Sovereign: imports nx_nofloat_autograd (pure integer) + nx_syscalls. 15// HONEST scope: this is the attention CORE (matmuls + causal softmax). RoPE + output-projection + multi-head are the 16// next sub-rung (CAP-NF-TRAIN-ATTN full). license_tier: ORIGINAL expect_exit: 0 17import "nx_nofloat_autograd.nx" 18import "nx_syscalls.nx" 19import "nx_gate_emit_lib.nx" 20 21const ALOG: *u8 = "knowledge/status/nofloat_attn.log" 22const Q16: i64 = 65536 23 24 25func g_abs(v: i64) -> i64 { if v < 0 { return 0 - v } return v } 26func q_milli(q: i64) -> i64 { var neg: i64=0; var a: i64=q; if a<0 { neg=1; a=0-a } let m: i64=(a*1000)/Q16; if neg==1 { return 0-m } return m } 27func a_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 } 28func a_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 } 29 30// ---- 2x2x2 matmul / matmul_nt gradcheck (op=8 matmul, op=9 matmul_nt) ---- 31func mm_loss(tape: *i64, vals: *i64, st: *i64, op: i64, As: *i64, Bs: *i64, Ts: *i64, leaves: *i64) -> i64 { 32 st[0]=0; st[1]=0 33 let nA: i64 = nfa_leaf(tape,vals,st,2,2,As,0) 34 let nB: i64 = nfa_leaf(tape,vals,st,2,2,Bs,0) 35 var nC: i64 = nA 36 if op == 8 { nC = nfa_matmul(tape,vals,st,nA,nB) } 37 if op == 9 { nC = nfa_matmul_nt(tape,vals,st,nA,nB) } 38 let nt: i64 = nfa_leaf(tape,vals,st,2,2,Ts,0) 39 let loss: i64 = nfa_mse(tape,vals,st,nC,nt) 40 leaves[0]=nA 41 return loss 42} 43func mm_lossval(tape: *i64, vals: *i64, st: *i64, op: i64, As: *i64, Bs: *i64, Ts: *i64) -> i64 { 44 let lv: *i64 = sys_mmap(8) as *i64 45 let loss: i64 = mm_loss(tape,vals,st,op,As,Bs,Ts,lv) 46 return nfa_val(tape,vals,loss,0) 47} 48func mm_fd(tape: *i64, vals: *i64, st: *i64, op: i64, As: *i64, Bs: *i64, Ts: *i64, pi: i64, h: i64) -> i64 { 49 let ap: *i64 = sys_mmap(4*8) as *i64 50 let am: *i64 = sys_mmap(4*8) as *i64 51 var i: i64=0 52 while i<4 { ap[i]=As[i]; am[i]=As[i]; i=i+1 } 53 ap[pi]=As[pi]+h; am[pi]=As[pi]-h 54 let lp: i64 = mm_lossval(tape,vals,st,op,ap,Bs,Ts) 55 let lm: i64 = mm_lossval(tape,vals,st,op,am,Bs,Ts) 56 return ((lp-lm)*Q16)/(2*h) 57} 58func mm_gradcheck(tape: *i64, vals: *i64, grads: *i64, st: *i64, op: i64, As: *i64, Bs: *i64, Ts: *i64, h: i64, tol_q: i64, floor_q: i64, worst: *i64) -> i64 { 59 let lv: *i64 = sys_mmap(8) as *i64 60 let loss: i64 = mm_loss(tape,vals,st,op,As,Bs,Ts,lv) 61 nfa_backward(tape,vals,grads,st[0],loss) 62 let nA: i64 = lv[0] 63 var ok: i64=1; worst[0]=0 64 var i: i64=0 65 while i<4 { 66 let ana: i64 = nfa_grad(tape,grads,nA,i) 67 let fd: i64 = mm_fd(tape,vals,st,op,As,Bs,Ts,i,h) 68 let num: i64 = g_abs(fd-ana) 69 var den: i64 = g_abs(ana); if den<floor_q { den=floor_q } 70 if num >= ((tol_q*den)>>16) { ok=0 } 71 let rel: i64 = (num*1000)/den 72 if rel>worst[0] { worst[0]=rel } 73 i=i+1 74 } 75 return ok 76} 77 78// ---- single-head causal attention forward on the tape; leaves[0..3]=nWq,nWk,nWv,nX; returns the O node ---- 79func attn_fwd(tape: *i64, vals: *i64, st: *i64, Xs: *i64, Wqs: *i64, Wks: *i64, Wvs: *i64, T: i64, d: i64, scale: i64, leaves: *i64) -> i64 { 80 st[0]=0; st[1]=0 81 let nX: i64 = nfa_leaf(tape,vals,st,T,d,Xs,0) 82 let nWq: i64 = nfa_leaf(tape,vals,st,d,d,Wqs,0) 83 let nWk: i64 = nfa_leaf(tape,vals,st,d,d,Wks,0) 84 let nWv: i64 = nfa_leaf(tape,vals,st,d,d,Wvs,0) 85 let nQ: i64 = nfa_matmul(tape,vals,st,nX,nWq) 86 let nK: i64 = nfa_matmul(tape,vals,st,nX,nWk) 87 let nV: i64 = nfa_matmul(tape,vals,st,nX,nWv) 88 let nS: i64 = nfa_matmul_nt(tape,vals,st,nQ,nK) 89 let nSs: i64 = nfa_cmul(tape,vals,st,nS,scale) 90 let nA: i64 = nfa_softmax_rows(tape,vals,st,nSs,1) 91 let nO: i64 = nfa_matmul(tape,vals,st,nA,nV) 92 leaves[0]=nWq; leaves[1]=nWk; leaves[2]=nWv; leaves[3]=nX 93 return nO 94} 95func attn_loss(tape: *i64, vals: *i64, st: *i64, Xs: *i64, Wqs: *i64, Wks: *i64, Wvs: *i64, Ts: *i64, T: i64, d: i64, scale: i64, leaves: *i64) -> i64 { 96 let nO: i64 = attn_fwd(tape,vals,st,Xs,Wqs,Wks,Wvs,T,d,scale,leaves) 97 let nt: i64 = nfa_leaf(tape,vals,st,T,d,Ts,0) 98 return nfa_mse(tape,vals,st,nO,nt) 99} 100func attn_lossval(tape: *i64, vals: *i64, st: *i64, Xs: *i64, Wqs: *i64, Wks: *i64, Wvs: *i64, Ts: *i64, T: i64, d: i64, scale: i64) -> i64 { 101 let lv: *i64 = sys_mmap(4*8) as *i64 102 let loss: i64 = attn_loss(tape,vals,st,Xs,Wqs,Wks,Wvs,Ts,T,d,scale,lv) 103 return nfa_val(tape,vals,loss,0) 104} 105// copy the O (attention output) values for given weights into outO (for realizable target generation) 106func attn_out(tape: *i64, vals: *i64, st: *i64, Xs: *i64, Wqs: *i64, Wks: *i64, Wvs: *i64, T: i64, d: i64, scale: i64, outO: *i64) -> i64 { 107 let lv: *i64 = sys_mmap(4*8) as *i64 108 let nO: i64 = attn_fwd(tape,vals,st,Xs,Wqs,Wks,Wvs,T,d,scale,lv) 109 var i: i64 = 0 110 while i < T*d { outO[i] = nfa_val(tape,vals,nO,i); i = i + 1 } 111 return 0 112} 113 114func main() -> i64 { 115 g_puts("nx_nofloat_attn gate (single-head CAUSAL attention backprops in PURE INTEGER Q16 -- MEASURED)\n" as *u8) 116 var pass: i64=0; var total: i64=0 117 let tape: *i64 = sys_mmap(512*7*8) as *i64 118 let vals: *i64 = sys_mmap(8192*8) as *i64 119 let grads: *i64 = sys_mmap(8192*8) as *i64 120 let st: *i64 = sys_mmap(2*8) as *i64 121 let h: i64 = 512 122 let floor_q: i64 = 4096 123 let worst: *i64 = sys_mmap(8) as *i64 124 125 // ---- A1: matmul gradcheck ---- 126 let A1: *i64 = sys_mmap(4*8) as *i64; A1[0]=32768; A1[1]=0-16384; A1[2]=49152; A1[3]=65536 127 let B1: *i64 = sys_mmap(4*8) as *i64; B1[0]=65536; B1[1]=16384; B1[2]=0-32768; B1[3]=49152 128 let T1: *i64 = sys_mmap(4*8) as *i64; T1[0]=13107; T1[1]=6554; T1[2]=0-19661; T1[3]=26214 129 let mm_ok: i64 = mm_gradcheck(tape,vals,grads,st,8,A1,B1,T1,h,4096,floor_q,worst) 130 g_puts(" [measure] matmul worst rel grad err = " as *u8); g_pn(worst[0]); g_puts(" /1000 (tol=62)\n" as *u8) 131 pass=pass+g_check("A1: matmul gradcheck -- C=A.B backward (dA=dC.B^T) == finite differences" as *u8, mm_ok); total=total+1 132 133 // ---- A2: matmul_nt gradcheck ---- 134 let mmnt_ok: i64 = mm_gradcheck(tape,vals,grads,st,9,A1,B1,T1,h,4096,floor_q,worst) 135 g_puts(" [measure] matmul_nt worst rel grad err = " as *u8); g_pn(worst[0]); g_puts(" /1000 (tol=62)\n" as *u8) 136 pass=pass+g_check("A2: matmul_nt gradcheck -- S=A.B^T (Q.K^T) backward == finite differences" as *u8, mmnt_ok); total=total+1 137 138 // ---- A3: full attention-core gradcheck wrt Wq (through causal softmax + both matmuls) ---- 139 let T: i64 = 3; let d: i64 = 2; let scale: i64 = 46341 // 1/sqrt(2) Q16 140 let Xs: *i64 = sys_mmap(T*d*8) as *i64; Xs[0]=32768; Xs[1]=16384; Xs[2]=0-16384; Xs[3]=49152; Xs[4]=65536; Xs[5]=0-32768 141 let Wq: *i64 = sys_mmap(d*d*8) as *i64; Wq[0]=49152; Wq[1]=0-16384; Wq[2]=32768; Wq[3]=65536 142 let Wk: *i64 = sys_mmap(d*d*8) as *i64; Wk[0]=16384; Wk[1]=32768; Wk[2]=0-32768; Wk[3]=49152 143 let Wv: *i64 = sys_mmap(d*d*8) as *i64; Wv[0]=65536; Wv[1]=0-32768; Wv[2]=16384; Wv[3]=49152 144 let Tg: *i64 = sys_mmap(T*d*8) as *i64; Tg[0]=13107; Tg[1]=0-6554; Tg[2]=19661; Tg[3]=6554; Tg[4]=0-13107; Tg[5]=26214 145 let lv: *i64 = sys_mmap(4*8) as *i64 146 let loss: i64 = attn_loss(tape,vals,st,Xs,Wq,Wk,Wv,Tg,T,d,scale,lv) 147 nfa_backward(tape,vals,grads,st[0],loss) 148 let nWq: i64 = lv[0] 149 var a3_ok: i64 = 1; var a3_worst: i64 = 0 150 var pi: i64 = 0 151 while pi < d*d { 152 let ana: i64 = nfa_grad(tape,grads,nWq,pi) 153 let wqp: *i64 = sys_mmap(d*d*8) as *i64 154 let wqm: *i64 = sys_mmap(d*d*8) as *i64 155 var z: i64=0 156 while z<d*d { wqp[z]=Wq[z]; wqm[z]=Wq[z]; z=z+1 } 157 wqp[pi]=Wq[pi]+h; wqm[pi]=Wq[pi]-h 158 let lp: i64 = attn_lossval(tape,vals,st,Xs,wqp,Wk,Wv,Tg,T,d,scale) 159 let lm: i64 = attn_lossval(tape,vals,st,Xs,wqm,Wk,Wv,Tg,T,d,scale) 160 let fd: i64 = ((lp-lm)*Q16)/(2*h) 161 let num: i64 = g_abs(fd-ana) 162 var den: i64 = g_abs(ana); if den<floor_q { den=floor_q } 163 if num >= ((8192*den)>>16) { a3_ok=0 } // tol 1/8 (long fixed-point chain through softmax) 164 let rel: i64 = (num*1000)/den 165 if rel>a3_worst { a3_worst=rel } 166 pi=pi+1 167 } 168 g_puts(" [measure] attention dL/dWq worst rel grad err = " as *u8); g_pn(a3_worst); g_puts(" /1000 (tol=125)\n" as *u8) 169 pass=pass+g_check("A3: attention-core gradcheck -- dL/dWq through causal-softmax + Q.K^T + A.V == finite diff" as *u8, a3_ok); total=total+1 170 171 // ---- D: neg-control teeth (matmul) ---- 172 let lvd: *i64 = sys_mmap(8) as *i64 173 let lossd: i64 = mm_loss(tape,vals,st,8,A1,B1,T1,lvd) 174 nfa_backward(tape,vals,grads,st[0],lossd) 175 let ana0: i64 = nfa_grad(tape,grads,lvd[0],0) 176 let fd0: i64 = mm_fd(tape,vals,st,8,A1,B1,T1,0,h) 177 let bad: i64 = 0 - ana0 178 var den0: i64 = g_abs(ana0); if den0<floor_q { den0=floor_q } 179 var caught: i64 = 1 180 if g_abs(fd0-bad) < ((4096*den0)>>16) { caught=0 } 181 pass=pass+g_check("D: neg-control -- a deliberately WRONG matmul grad is rejected (teeth)" as *u8, caught); total=total+1 182 183 // ---- B: attention value-path LEARNS (Wq,Wk fixed -> A fixed -> O=A.(X.Wv) convex in Wv) ---- 184 let Wvt: *i64 = sys_mmap(d*d*8) as *i64; Wvt[0]=65536; Wvt[1]=0-32768; Wvt[2]=16384; Wvt[3]=49152 // Wv* target 185 let tgt: *i64 = sys_mmap(T*d*8) as *i64 186 attn_out(tape,vals,st,Xs,Wq,Wk,Wvt,T,d,scale,tgt) // realizable target = attention out with Wv* 187 let Wvp: *i64 = sys_mmap(d*d*8) as *i64; Wvp[0]=0; Wvp[1]=0; Wvp[2]=0; Wvp[3]=0 // learn from zero 188 let lvb: *i64 = sys_mmap(4*8) as *i64 189 let gW: *i64 = sys_mmap(d*d*8) as *i64 190 var lf: i64 = 0; var ll: i64 = 0 191 var ep: i64 = 0 192 while ep < 6000 { 193 let lossb: i64 = attn_loss(tape,vals,st,Xs,Wq,Wk,Wvp,tgt,T,d,scale,lvb) 194 nfa_backward(tape,vals,grads,st[0],lossb) 195 if ep==0 { lf = nfa_val(tape,vals,lossb,0) } 196 ll = nfa_val(tape,vals,lossb,0) 197 var z: i64=0 198 while z<d*d { gW[z]=nfa_grad(tape,grads,lvb[2],z); z=z+1 } // lvb[2] = nWv leaf 199 nfa_sgd(Wvp, gW, d*d, 1024) 200 ep=ep+1 201 } 202 g_puts(" [measure] attention Wv-learn loss: start=" as *u8); g_pn(lf); g_puts(" end=" as *u8); g_pn(ll) 203 g_puts(" Wv=[" as *u8); g_pn(Wvp[0]); g_puts("," as *u8); g_pn(Wvp[1]); g_puts("," as *u8); g_pn(Wvp[2]); g_puts("," as *u8); g_pn(Wvp[3]); g_puts("] vs Wv*=[65536,-32768,16384,49152]\n" as *u8) 204 var learns: i64 = 1 205 if ll*10 > lf { learns=0 } 206 if g_abs(Wvp[0]-65536)>9830 { learns=0 } 207 if g_abs(Wvp[1]+32768)>9830 { learns=0 } 208 if g_abs(Wvp[2]-16384)>9830 { learns=0 } 209 if g_abs(Wvp[3]-49152)>9830 { learns=0 } 210 if lf<=0 { learns=0 } 211 pass=pass+g_check("B: attention value-path LEARNS in pure integer -- Wv converges to target, loss collapses" as *u8, learns); total=total+1 212 213 // ---- C: bit-exact ---- 214 let Wvp2: *i64 = sys_mmap(d*d*8) as *i64; Wvp2[0]=0; Wvp2[1]=0; Wvp2[2]=0; Wvp2[3]=0 215 let lvc: *i64 = sys_mmap(4*8) as *i64 216 let gW2: *i64 = sys_mmap(d*d*8) as *i64 217 var ep2: i64=0 218 while ep2 < 6000 { 219 let lossc: i64 = attn_loss(tape,vals,st,Xs,Wq,Wk,Wvp2,tgt,T,d,scale,lvc) 220 nfa_backward(tape,vals,grads,st[0],lossc) 221 var z: i64=0 222 while z<d*d { gW2[z]=nfa_grad(tape,grads,lvc[2],z); z=z+1 } 223 nfa_sgd(Wvp2, gW2, d*d, 1024) 224 ep2=ep2+1 225 } 226 var bitexact: i64 = 1 227 var z2: i64=0 228 while z2<d*d { if Wvp2[z2]!=Wvp[z2] { bitexact=0 } z2=z2+1 } 229 pass=pass+g_check("C: bit-exact -- training twice gives IDENTICAL integer Wv (determinism)" as *u8, bitexact); total=total+1 230 231 // ---- emit ---- 232 var okall: i64=0; if pass==total { okall=1 } 233 let logf: i64 = sys_openat_append(ALOG, 420) 234 if logf >= 0 { 235 a_ws(logf,"NOFLOATATTN ops=matmul,matmul_nt,cmul,causal_softmax_rows A1=" as *u8); a_wn(logf,mm_ok) 236 a_ws(logf," A2=" as *u8); a_wn(logf,mmnt_ok); a_ws(logf," A3_attn=" as *u8); a_wn(logf,a3_ok); a_ws(logf," D=" as *u8); a_wn(logf,caught) 237 a_ws(logf," B_learns=" as *u8); a_wn(logf,learns); a_ws(logf," C_bitexact=" as *u8); a_wn(logf,bitexact) 238 if okall==1 { a_ws(logf," verdict=GREEN\n" as *u8) } else { a_ws(logf," verdict=RED\n" as *u8) } 239 sys_close(logf) 240 } 241 g_puts("---- nofloat_attn gate: passed " as *u8); g_pn(pass); g_puts(" / " as *u8); g_pn(total); g_puts(" ----\n" as *u8) 242 if okall==1 { g_puts("verdict=GREEN\n" as *u8); sys_exit(0); return 0 } 243 g_puts("verdict=RED\n" as *u8); sys_exit(1); return 1 244}