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}