code wiki / _hdl_build / nx_intfp_swiglu_gradcheck_gate.nx

nx_intfp_swiglu_gradcheck_gate.nx source

↩ module page · 100 lines · 6973 B

1// nx_intfp_swiglu_gradcheck_gate.nx -- the OTHER transformer sublayer in integer: SwiGLU FFN, gradchecked. 2// y = (SiLU(x@Wg) (.) (x@Wu)) @ Wd , SiLU(z)=z*sigmoid(z)=z/(1+exp(-z)) -- reuses the proven fixed-point exp. 3// Full backward (dy->dWd, dm-> da/du -> dg via SiLU', -> dWg/dWu), ENTIRELY Q20 integer, gradchecked (integer 4// finite-diff, max-|grad|-relative metric) on all 3 weight matrices. SiLU'(z)=sig(z)(1+z(1-sig(z))). With this + 5// attention(48/48) both sublayers of a transformer layer are proven in integer. No float. license_tier: ORIGINAL 6import "nx_syscalls.nx" 7 8func w(s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} sys_write(1,s,n); return 0 } 9func wn(v: i64) -> i64 { if v==0 { sys_write(1,"0" as *u8,1); return 0 } var m: i64=v; if m<0{sys_write(1,"-" as *u8,1);m=0-m} let t: *u8=sys_mmap(24); var k: i64=0; while m>0{t[k]=(48+(m%10)) as u8;m=m/10;k=k+1} let o: *u8=sys_mmap(24); var q: i64=k-1; var i: i64=0; while q>=0{o[i]=t[q];i=i+1;q=q-1} sys_write(1,o,i); return 0 } 10func iabs(v: i64) -> i64 { if v<0 { return 0-v } return v } 11 12const S: i64 = 1048576 // Q20 13const T: i64 = 3 14const DM: i64 = 4 15const HF: i64 = 6 16 17func fp_exp(xq: i64) -> i64 { 18 let y: i64=(xq*1512776)/S 19 var yi: i64=0 20 if y>=0 { yi=y/S } else { yi=0-(((0-y)+S-1)/S) } 21 let yf: i64=y-yi*S 22 var p: i64=10085 23 p=58197+(p*yf)/S; p=251882+(p*yf)/S; p=726817+(p*yf)/S; p=S+(p*yf)/S 24 if yi>=0 { if yi>=31 { return 2000000000 } return p*(1<<yi) } 25 let k: i64=0-yi; if k>=31 { return 0 } 26 return p/(1<<k) 27} 28func sigmoid(z: i64) -> i64 { let e: i64=fp_exp(0-z); return (S*S)/(S+e) } // 1/(1+exp(-z)), Q20 29 30// forward: fills g,a,u,m,y; returns L_q32 31func swiglu_fwd(x: *i64, Wg: *i64, Wu: *i64, Wd: *i64, g: *i64, a: *i64, u: *i64, m: *i64, y: *i64) -> i64 { 32 var t: i64=0 33 while t<T { var j: i64=0; while j<HF { var gg: i64=0; var uu: i64=0; var k: i64=0 34 while k<DM { gg=gg+x[t*DM+k]*Wg[k*HF+j]; uu=uu+x[t*DM+k]*Wu[k*HF+j]; k=k+1 } 35 let gv: i64=gg/S; g[t*HF+j]=gv; u[t*HF+j]=uu/S 36 let sig: i64=sigmoid(gv); let av: i64=(gv*sig)/S; a[t*HF+j]=av; m[t*HF+j]=(av*u[t*HF+j])/S; j=j+1 } t=t+1 } 37 var Lq: i64=0; t=0 38 while t<T { var i: i64=0; while i<DM { var acc: i64=0; var j: i64=0; while j<HF { acc=acc+m[t*HF+j]*Wd[j*DM+i]; j=j+1 } let yv: i64=acc/S; y[t*DM+i]=yv; Lq=Lq+yv*yv; i=i+1 } t=t+1 } 39 return Lq 40} 41 42// backward (single fn; Wd needed for dm) 43func swiglu_bwd2(x: *i64, Wd: *i64, g: *i64, a: *i64, u: *i64, m: *i64, y: *i64, dWg: *i64, dWu: *i64, dWd: *i64) -> i64 { 44 let dy: *i64=sys_mmap(T*DM*8) as *i64; let dm: *i64=sys_mmap(T*HF*8) as *i64 45 let dg: *i64=sys_mmap(T*HF*8) as *i64; let du: *i64=sys_mmap(T*HF*8) as *i64 46 var t: i64=0; while t<T { var i: i64=0; while i<DM { dy[t*DM+i]=2*y[t*DM+i]; i=i+1 } t=t+1 } 47 var j: i64=0; while j<HF { var i: i64=0; while i<DM { var acc: i64=0; t=0; while t<T { acc=acc+(m[t*HF+j]*dy[t*DM+i])/S; t=t+1 } dWd[j*DM+i]=acc; i=i+1 } j=j+1 } 48 // dm[t,j] = Σ_i dy[t,i] Wd[j,i] 49 t=0; while t<T { j=0; while j<HF { var acc: i64=0; var i: i64=0; while i<DM { acc=acc+(dy[t*DM+i]*Wd[j*DM+i])/S; i=i+1 } dm[t*HF+j]=acc; j=j+1 } t=t+1 } 50 // da=dm⊙u ; du=dm⊙a ; dg = da * SiLU'(g), SiLU'(z)=sig(1+z(1-sig)) 51 t=0; while t<T { j=0; while j<HF { let dmv: i64=dm[t*HF+j]; let dav: i64=(dmv*u[t*HF+j])/S; du[t*HF+j]=(dmv*a[t*HF+j])/S 52 let z: i64=g[t*HF+j]; let sig: i64=sigmoid(z); let zt: i64=(z*(S-sig))/S; let dsil: i64=(sig*(S+zt))/S; dg[t*HF+j]=(dav*dsil)/S; j=j+1 } t=t+1 } 53 // dWg[k,j]=Σ_t x[t,k] dg[t,j] ; dWu[k,j]=Σ_t x[t,k] du[t,j] 54 var k: i64=0; while k<DM { j=0; while j<HF { var ag: i64=0; var au: i64=0; t=0; while t<T { ag=ag+(x[t*DM+k]*dg[t*HF+j])/S; au=au+(x[t*DM+k]*du[t*HF+j])/S; t=t+1 } dWg[k*HF+j]=ag; dWu[k*HF+j]=au; j=j+1 } k=k+1 } 55 return 0 56} 57 58func gcheck(name: *u8, x: *i64, Wg: *i64, Wu: *i64, Wd: *i64, Wtgt: *i64, dW: *i64, ncell: i64, g: *i64, a: *i64, u: *i64, m: *i64, y: *i64) -> i64 { 59 let DELTA: i64=10486; let TOLP: i64=60 60 var maxabs: i64=1; var q: i64=0; while q<ncell { if iabs(dW[q])>maxabs { maxabs=iabs(dW[q]) } q=q+1 } 61 var npass: i64=0; var worst: i64=0; var i: i64=0 62 while i<ncell { 63 let save: i64=Wtgt[i] 64 Wtgt[i]=save+DELTA; let Lp: i64=swiglu_fwd(x,Wg,Wu,Wd,g,a,u,m,y) 65 Wtgt[i]=save-DELTA; let Lm: i64=swiglu_fwd(x,Wg,Wu,Wd,g,a,u,m,y) 66 Wtgt[i]=save; let dd: i64=swiglu_fwd(x,Wg,Wu,Wd,g,a,u,m,y) 67 let num: i64=(Lp-Lm)/(2*DELTA); let rel: i64=(iabs(num-dW[i])*1000)/maxabs 68 if rel<=TOLP { npass=npass+1 } else { w(" " as *u8); w(name); w(" cell " as *u8); wn(i); w(" ana=" as *u8); wn(dW[i]); w(" num=" as *u8); wn(num); w(" rel/max=" as *u8); wn(rel); w("\n" as *u8) } 69 if rel>worst { worst=rel } 70 i=i+1 71 } 72 w(" " as *u8); w(name); w(": " as *u8); wn(npass); w("/" as *u8); wn(ncell); w(" (err<=6% of max|grad|=" as *u8); wn(maxabs); w("), worst=" as *u8); wn(worst); w("permil\n" as *u8) 73 return npass 74} 75 76func main() -> i64 { 77 w("=== nx_intfp_swiglu_gradcheck: Q20 SwiGLU FFN (SiLU=z*sigmoid(z)) full backward, gradcheck Wg/Wu/Wd -- no float ===\n\n" as *u8) 78 let x: *i64=sys_mmap(T*DM*8) as *i64 79 let Wg: *i64=sys_mmap(DM*HF*8) as *i64; let Wu: *i64=sys_mmap(DM*HF*8) as *i64; let Wd: *i64=sys_mmap(HF*DM*8) as *i64 80 let g: *i64=sys_mmap(T*HF*8) as *i64; let a: *i64=sys_mmap(T*HF*8) as *i64; let u: *i64=sys_mmap(T*HF*8) as *i64 81 let m: *i64=sys_mmap(T*HF*8) as *i64; let y: *i64=sys_mmap(T*DM*8) as *i64 82 let dWg: *i64=sys_mmap(DM*HF*8) as *i64; let dWu: *i64=sys_mmap(DM*HF*8) as *i64; let dWd: *i64=sys_mmap(HF*DM*8) as *i64 83 var i: i64=0; while i<T*DM { x[i]=((((i*5+2)%11)-5)*S)/10; i=i+1 } 84 i=0; while i<DM*HF { Wg[i]=((((i*7+1)%13)-6)*S)/14; Wu[i]=((((i*3+4)%13)-6)*S)/14; i=i+1 } 85 i=0; while i<HF*DM { Wd[i]=((((i*5+3)%13)-6)*S)/14; i=i+1 } 86 87 let L0: i64=swiglu_fwd(x,Wg,Wu,Wd,g,a,u,m,y) 88 swiglu_bwd2(x,Wd,g,a,u,m,y,dWg,dWu,dWd) 89 w(" forward L_q32=" as *u8); wn(L0); w(" (SiLU=z*sigmoid(z), reuses fp_exp; T=" as *u8); wn(T); w(" d=" as *u8); wn(DM); w(" h=" as *u8); wn(HF); w(")\n" as *u8) 90 w(" gradcheck (only failing cells printed):\n" as *u8) 91 let pd: i64=gcheck("Wd" as *u8, x, Wg, Wu, Wd, Wd, dWd, HF*DM, g,a,u,m,y) 92 let pu: i64=gcheck("Wu" as *u8, x, Wg, Wu, Wd, Wu, dWu, DM*HF, g,a,u,m,y) 93 let pg: i64=gcheck("Wg" as *u8, x, Wg, Wu, Wd, Wg, dWg, DM*HF, g,a,u,m,y) // thru SiLU' 94 let tot: i64=pd+pu+pg; let want: i64=HF*DM+DM*HF+DM*HF 95 w("\n SwiGLU FFN composition gradcheck: " as *u8); wn(tot); w("/" as *u8); wn(want); w(" cells correct\n" as *u8) 96 w("NX-INTFP-SWIGLU-GRADCHECK verdict=" as *u8) 97 if tot==want { w("GREEN " as *u8); wn(tot); w("/" as *u8); wn(want); w(" -- integer SwiGLU FFN + SiLU proven. Both transformer sublayers (attn 48/48 + FFN) now integer.\n" as *u8) } 98 else { w("RED " as *u8); wn(tot); w("/" as *u8); wn(want); w(" -- SiLU/FFN Q-scaling bug (see cells)\n" as *u8) } 99 return 0 100}