code wiki / _hdl_build / nx_intfp_block_gradcheck_gate.nx

nx_intfp_block_gradcheck_gate.nx source

↩ module page · 165 lines · 12300 B

1// nx_intfp_block_gradcheck_gate.nx -- CAPSTONE assembly pattern: a pre-norm residual block y = x + Attn(RMSNorm(x)) 2// with the FULL end-to-end integer backward, validating the two NEW assembly risks the isolated gates didn't cover: 3// (1) INPUT-gradient propagation dX through a sublayer (attention dh), and 4// (2) RESIDUAL gradient accumulation (dx = dx_residual + dx_norm-path). 5// Gradchecked (integer finite-diff, max-|grad| metric) on gamma, Wq, Wv AND the input X (which flows through BOTH 6// the residual and the norm->attention path), then TRAINED (loss drops). If this holds, the full transformer layer 7// is the same pattern twice (attn block + FFN block). Q20 integer, no float. Ptr-array plumbing to fit arg budget. 8// slots: 0x 1gamma 2Wq 3Wk 4Wv 5Y | 6h 7Q 8K 9V 10A 11O 12y 13rms | 14dgamma 15dWq 16dWk 17dWv 18dX 9// license_tier: ORIGINAL 10import "nx_syscalls.nx" 11 12func w(s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} sys_write(1,s,n); return 0 } 13func 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 } 14func iabs(v: i64) -> i64 { if v<0 { return 0-v } return v } 15 16const S: i64 = 1048576 // Q20 17const T: i64 = 3 18const DM: i64 = 4 19const SCALE: i64 = 524288 // 1/sqrt(4)=0.5 20const EPS: i64 = 1048576 // eps in Q40 (tiny vs ms) 21 22func isqrt(n: i64) -> i64 { 23 if n<=0 { return 0 } 24 var bit: i64=1; while bit*4<=n { bit=bit*4 } 25 var res: i64=0; var num: i64=n 26 while bit!=0 { if num>=res+bit { num=num-(res+bit); res=(res/2)+bit } else { res=res/2 } bit=bit/4 } 27 return res 28} 29func fp_exp(xq: i64) -> i64 { 30 let y: i64=(xq*1512776)/S 31 var yi: i64=0 32 if y>=0 { yi=y/S } else { yi=0-(((0-y)+S-1)/S) } 33 let yf: i64=y-yi*S 34 var p: i64=10085 35 p=58197+(p*yf)/S; p=251882+(p*yf)/S; p=726817+(p*yf)/S; p=S+(p*yf)/S 36 if yi>=0 { if yi>=31 { return 2000000000 } return p*(1<<yi) } 37 let k: i64=0-yi; if k>=31 { return 0 } 38 return p/(1<<k) 39} 40 41// forward: fills h,Q,K,V,A,O,y,rms; returns L_q32 = sum((y-Y)^2) 42func block_fwd(P: *i64) -> i64 { 43 let X: *i64=P[0] as *i64; let gm: *i64=P[1] as *i64; let Wq: *i64=P[2] as *i64; let Wk: *i64=P[3] as *i64; let Wv: *i64=P[4] as *i64; let Y: *i64=P[5] as *i64 44 let h: *i64=P[6] as *i64; let Q: *i64=P[7] as *i64; let K: *i64=P[8] as *i64; let V: *i64=P[9] as *i64; let A: *i64=P[10] as *i64; let O: *i64=P[11] as *i64; let y: *i64=P[12] as *i64; let rms: *i64=P[13] as *i64 45 // RMSNorm1: h = (X/rms)*gamma 46 var t: i64=0 47 while t<T { var ms: i64=0; var i: i64=0; while i<DM { ms=ms+X[t*DM+i]*X[t*DM+i]; i=i+1 } ms=ms/DM+EPS; let r: i64=isqrt(ms); rms[t]=r; let inv: i64=(S*S)/r 48 i=0; while i<DM { let nrm: i64=(X[t*DM+i]*inv)/S; h[t*DM+i]=(nrm*gm[i])/S; i=i+1 } t=t+1 } 49 // Attention(h) 50 t=0 51 while t<T { var i: i64=0; while i<DM { var aq: i64=0; var ak: i64=0; var av: i64=0; var k: i64=0 52 while k<DM { aq=aq+h[t*DM+k]*Wq[k*DM+i]; ak=ak+h[t*DM+k]*Wk[k*DM+i]; av=av+h[t*DM+k]*Wv[k*DM+i]; k=k+1 } 53 Q[t*DM+i]=aq/S; K[t*DM+i]=ak/S; V[t*DM+i]=av/S; i=i+1 } t=t+1 } 54 t=0 55 while t<T { var mx: i64=0-2000000000; var s: i64=0 56 while s<=t { var dot: i64=0; var i: i64=0; while i<DM { dot=dot+Q[t*DM+i]*K[s*DM+i]; i=i+1 } let sc: i64=((dot/S)*SCALE)/S; A[t*T+s]=sc; if sc>mx { mx=sc } s=s+1 } 57 var sum: i64=0; s=0; while s<=t { let e: i64=fp_exp(A[t*T+s]-mx); A[t*T+s]=e; sum=sum+e; s=s+1 } 58 s=0; while s<=t { A[t*T+s]=(A[t*T+s]*S+sum/2)/sum; s=s+1 } t=t+1 } 59 // O = A@V ; residual y = X + O ; loss 60 var Lq: i64=0; t=0 61 while t<T { var i: i64=0; while i<DM { var acc: i64=0; var s: i64=0; while s<=t { acc=acc+A[t*T+s]*V[s*DM+i]; s=s+1 } O[t*DM+i]=acc/S 62 let yv: i64=X[t*DM+i]+O[t*DM+i]; y[t*DM+i]=yv; let e: i64=yv-Y[t*DM+i]; Lq=Lq+e*e; i=i+1 } t=t+1 } 63 return Lq 64} 65 66// backward: fills dgamma,dWq,dWk,dWv,dX 67func block_bwd(P: *i64) -> i64 { 68 let X: *i64=P[0] as *i64; let gm: *i64=P[1] as *i64; let Wq: *i64=P[2] as *i64; let Wk: *i64=P[3] as *i64; let Wv: *i64=P[4] as *i64; let Y: *i64=P[5] as *i64 69 let h: *i64=P[6] as *i64; let Q: *i64=P[7] as *i64; let K: *i64=P[8] as *i64; let V: *i64=P[9] as *i64; let A: *i64=P[10] as *i64; let O: *i64=P[11] as *i64; let y: *i64=P[12] as *i64; let rms: *i64=P[13] as *i64 70 let dgm: *i64=P[14] as *i64; let dWq: *i64=P[15] as *i64; let dWk: *i64=P[16] as *i64; let dWv: *i64=P[17] as *i64; let dX: *i64=P[18] as *i64 71 let dO: *i64=sys_mmap(T*DM*8) as *i64; let dV: *i64=sys_mmap(T*DM*8) as *i64; let dA: *i64=sys_mmap(T*T*8) as *i64; let dsc: *i64=sys_mmap(T*T*8) as *i64 72 let dQ: *i64=sys_mmap(T*DM*8) as *i64; let dK: *i64=sys_mmap(T*DM*8) as *i64; let dh: *i64=sys_mmap(T*DM*8) as *i64 73 // dy = 2(y-Y) ; residual: dO=dy, dX starts = dy 74 var t: i64=0; while t<T { var i: i64=0; while i<DM { let dy: i64=2*(y[t*DM+i]-Y[t*DM+i]); dO[t*DM+i]=dy; dX[t*DM+i]=dy; i=i+1 } t=t+1 } 75 // attention backward (input h). dV,dA,dsc,dQ,dK 76 var s: i64=0; while s<DM*T { dV[s]=0; s=s+1 } 77 s=0; while s<T { var i: i64=0; while i<DM { var acc: i64=0; t=s; while t<T { acc=acc+(A[t*T+s]*dO[t*DM+i])/S; t=t+1 } dV[s*DM+i]=acc; i=i+1 } s=s+1 } 78 t=0; while t<T { s=0; while s<=t { var acc: i64=0; var i: i64=0; while i<DM { acc=acc+(dO[t*DM+i]*V[s*DM+i])/S; i=i+1 } dA[t*T+s]=acc; s=s+1 } t=t+1 } 79 t=0; while t<T { var dot: i64=0; s=0; while s<=t { dot=dot+(A[t*T+s]*dA[t*T+s])/S; s=s+1 } s=0; while s<=t { dsc[t*T+s]=(A[t*T+s]*(dA[t*T+s]-dot))/S; s=s+1 } t=t+1 } 80 t=0; while t<T*DM { dQ[t]=0; dK[t]=0; t=t+1 } 81 t=0; while t<T { var i: i64=0; while i<DM { var acc: i64=0; s=0; while s<=t { let dqk: i64=(SCALE*dsc[t*T+s])/S; acc=acc+(dqk*K[s*DM+i])/S; s=s+1 } dQ[t*DM+i]=acc; i=i+1 } t=t+1 } 82 s=0; while s<T { var i: i64=0; while i<DM { var acc: i64=0; t=s; while t<T { let dqk: i64=(SCALE*dsc[t*T+s])/S; acc=acc+(dqk*Q[t*DM+i])/S; t=t+1 } dK[s*DM+i]=acc; i=i+1 } s=s+1 } 83 // dW (from h) and dh (input grad of attention) = sum over the three projections 84 var kk: i64=0 85 while kk<DM { var i: i64=0; while i<DM { var aq: i64=0; var ak: i64=0; var av: i64=0; t=0; while t<T { aq=aq+(h[t*DM+kk]*dQ[t*DM+i])/S; ak=ak+(h[t*DM+kk]*dK[t*DM+i])/S; av=av+(h[t*DM+kk]*dV[t*DM+i])/S; t=t+1 } dWq[kk*DM+i]=aq; dWk[kk*DM+i]=ak; dWv[kk*DM+i]=av; i=i+1 } kk=kk+1 } 86 t=0; while t<T { var k2: i64=0; while k2<DM { var acc: i64=0; var i: i64=0; while i<DM { acc=acc+(dQ[t*DM+i]*Wq[k2*DM+i])/S + (dK[t*DM+i]*Wk[k2*DM+i])/S + (dV[t*DM+i]*Wv[k2*DM+i])/S; i=i+1 } dh[t*DM+k2]=acc; k2=k2+1 } t=t+1 } 87 // RMSNorm backward: dh -> (dgamma, and dX += norm-path). norm=X/rms, per row t. 88 var i2: i64=0; while i2<DM { dgm[i2]=0; i2=i2+1 } 89 t=0 90 while t<T { 91 let r: i64=rms[t]; let inv: i64=(S*S)/r // 1/rms Q20 92 let invr3: i64=(((inv*inv)/S)*inv)/S // inv^3 93 // dn_i = dh_i * gamma_i ; dgamma_i += dh_i * norm_i(=X_i/rms) 94 var c: i64=0; var i: i64=0 95 while i<DM { let nrm: i64=(X[t*DM+i]*inv)/S; dgm[i]=dgm[i]+(dh[t*DM+i]*nrm)/S; let dn: i64=(dh[t*DM+i]*gm[i])/S; let ggx: i64=(dn*X[t*DM+i])/S; c=c+ggx; i=i+1 } // c=sum dn_i X_i 96 i=0 97 while i<DM { let dn: i64=(dh[t*DM+i]*gm[i])/S; let term1: i64=(dn*inv)/S; let tt: i64=(X[t*DM+i]*c)/S; let term2: i64=(((tt*invr3)/S))/DM; dX[t*DM+i]=dX[t*DM+i]+(term1-term2); i=i+1 } 98 t=t+1 99 } 100 return 0 101} 102 103func gcheck(name: *u8, P: *i64, tgtslot: i64, dW: *i64, ncell: i64) -> i64 { 104 let Wtgt: *i64=P[tgtslot] as *i64 105 let DELTA: i64=10486; let TOLP: i64=60 106 var maxabs: i64=1; var q: i64=0; while q<ncell { if iabs(dW[q])>maxabs { maxabs=iabs(dW[q]) } q=q+1 } 107 var npass: i64=0; var worst: i64=0; var i: i64=0 108 while i<ncell { 109 let save: i64=Wtgt[i] 110 Wtgt[i]=save+DELTA; let Lp: i64=block_fwd(P) 111 Wtgt[i]=save-DELTA; let Lm: i64=block_fwd(P) 112 Wtgt[i]=save; let dd: i64=block_fwd(P) 113 let num: i64=(Lp-Lm)/(2*DELTA); let rel: i64=(iabs(num-dW[i])*1000)/maxabs 114 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=" as *u8); wn(rel); w("\n" as *u8) } 115 if rel>worst { worst=rel } 116 i=i+1 117 } 118 w(" " as *u8); w(name); w(": " as *u8); wn(npass); w("/" as *u8); wn(ncell); w(" (worst=" as *u8); wn(worst); w("permil of max|grad|=" as *u8); wn(maxabs); w(")\n" as *u8) 119 return npass 120} 121 122func main() -> i64 { 123 w("=== nx_intfp_block_gradcheck: y = X + Attn(RMSNorm(X)) -- full assembled backward (dX + residual), Q20 integer ===\n\n" as *u8) 124 let P: *i64=sys_mmap(20*8) as *i64 125 // params 126 P[0]=sys_mmap(T*DM*8); P[1]=sys_mmap(DM*8); P[2]=sys_mmap(DM*DM*8); P[3]=sys_mmap(DM*DM*8); P[4]=sys_mmap(DM*DM*8); P[5]=sys_mmap(T*DM*8) 127 // fwd buffers 128 P[6]=sys_mmap(T*DM*8); P[7]=sys_mmap(T*DM*8); P[8]=sys_mmap(T*DM*8); P[9]=sys_mmap(T*DM*8); P[10]=sys_mmap(T*T*8); P[11]=sys_mmap(T*DM*8); P[12]=sys_mmap(T*DM*8); P[13]=sys_mmap(T*8) 129 // grads 130 P[14]=sys_mmap(DM*8); P[15]=sys_mmap(DM*DM*8); P[16]=sys_mmap(DM*DM*8); P[17]=sys_mmap(DM*DM*8); P[18]=sys_mmap(T*DM*8) 131 let X: *i64=P[0] as *i64; let gm: *i64=P[1] as *i64; let Wq: *i64=P[2] as *i64; let Wk: *i64=P[3] as *i64; let Wv: *i64=P[4] as *i64; let Y: *i64=P[5] as *i64 132 var i: i64=0; while i<T*DM { X[i]=((((i*5+2)%11)-5)*S)/10; i=i+1 } 133 i=0; while i<DM { gm[i]=S; i=i+1 } // gamma init = 1 134 i=0; while i<DM*DM { Wq[i]=((((i*7+1)%13)-6)*S)/16; Wk[i]=((((i*3+5)%13)-6)*S)/16; Wv[i]=((((i*11+2)%13)-6)*S)/16; i=i+1 } 135 i=0; while i<T*DM { Y[i]=((((i*3+1)%7)-3)*S)/8; i=i+1 } // regression target 136 137 let L0: i64=block_fwd(P); block_bwd(P) 138 w(" forward L_q32=" as *u8); wn(L0); w("\n gradcheck (assembled backward; only failing cells printed):\n" as *u8) 139 let pg: i64=gcheck("gamma" as *u8, P, 1, P[14] as *i64, DM) 140 let pv: i64=gcheck("Wv " as *u8, P, 4, P[17] as *i64, DM*DM) 141 let pq: i64=gcheck("Wq " as *u8, P, 2, P[15] as *i64, DM*DM) 142 let px: i64=gcheck("X(in)" as *u8, P, 0, P[18] as *i64, T*DM) // input grad thru BOTH residual + norm->attn 143 let tot: i64=pg+pv+pq+px; let want: i64=DM+DM*DM+DM*DM+T*DM 144 w("\n ASSEMBLED-BLOCK gradcheck: " as *u8); wn(tot); w("/" as *u8); wn(want); w(" cells correct\n" as *u8) 145 146 // TRAIN the assembled block by integer SGD (error-feedback) -> loss must drop 147 let lr: i64=(S*10)/100 148 let R1: *i64=sys_mmap(DM*8) as *i64; let R2: *i64=sys_mmap(DM*DM*8) as *i64; let R3: *i64=sys_mmap(DM*DM*8) as *i64; let R4: *i64=sys_mmap(DM*DM*8) as *i64 149 i=0; while i<DM { R1[i]=0; i=i+1 } i=0; while i<DM*DM { R2[i]=0; R3[i]=0; R4[i]=0; i=i+1 } 150 let dgm: *i64=P[14] as *i64; let dWq2: *i64=P[15] as *i64; let dWk2: *i64=P[16] as *i64; let dWv2: *i64=P[17] as *i64 151 let gm2: *i64=P[1] as *i64; let Wq2: *i64=P[2] as *i64; let Wk2: *i64=P[3] as *i64; let Wv2: *i64=P[4] as *i64 152 var step: i64=1; var Lprev: i64=L0 153 while step<=1500 { 154 let L: i64=block_fwd(P); block_bwd(P) 155 var a: i64=0; while a<DM { R1[a]=R1[a]+lr*dgm[a]; let tk: i64=R1[a]/S; gm2[a]=gm2[a]-tk; R1[a]=R1[a]-tk*S; a=a+1 } 156 a=0; while a<DM*DM { R2[a]=R2[a]+lr*dWq2[a]; let tk: i64=R2[a]/S; Wq2[a]=Wq2[a]-tk; R2[a]=R2[a]-tk*S; R3[a]=R3[a]+lr*dWk2[a]; let tk2: i64=R3[a]/S; Wk2[a]=Wk2[a]-tk2; R3[a]=R3[a]-tk2*S; R4[a]=R4[a]+lr*dWv2[a]; let tk3: i64=R4[a]/S; Wv2[a]=Wv2[a]-tk3; R4[a]=R4[a]-tk3*S; a=a+1 } 157 Lprev=L; step=step+1 158 } 159 let Lf: i64=block_fwd(P) 160 w(" TRAIN: loss_q32 " as *u8); wn(L0); w(" -> " as *u8); wn(Lf); w(" over 1500 steps\n" as *u8) 161 w("NX-INTFP-BLOCK verdict=" as *u8) 162 if tot==want { if Lf*4<L0 { w("GREEN gradcheck " as *u8); wn(tot); w("/" as *u8); wn(want); w(" + TRAINS (loss fell >4x) -- assembled block (dX + residual + norm+attn) proven; full layer = this pattern twice\n" as *u8) } else { w("YELLOW gradcheck passed but loss didn't fall >4x\n" as *u8) } } 163 else { w("RED gradcheck " as *u8); wn(tot); w("/" as *u8); wn(want); w(" -- assembly backward bug (dX/residual/norm chaining)\n" as *u8) } 164 return 0 165}