code wiki / _hdl_build / nx_intfp_attention_gradcheck_gate.nx
nx_intfp_attention_gradcheck_gate.nx source
↩ module page · 117 lines · 8424 B
1// nx_intfp_attention_gradcheck_gate.nx -- THE COMPOSITION PROOF: a full single-head CAUSAL self-attention block
2// (Q,K,V = X@W ; scores = scale*QK^T ; A = softmax(scores) causal ; O = A@V ; L = sum(O^2)) with the COMPLETE
3// backward, done ENTIRELY in Q16 integer, gradchecked (integer finite-diff) on all three projection matrices
4// Wq/Wk/Wv. Wq/Wk gradients flow through the softmax Jacobian AND two matmul-transpose backprops -- if THESE
5// pass, the integer ops compose correctly and the full transformer tape is assemblable. No float anywhere.
6// scale = 1/sqrt(d); d=4 -> scale=0.5 exactly (Q16 32768). license_tier: ORIGINAL
7import "nx_syscalls.nx"
8
9func w(s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} sys_write(1,s,n); return 0 }
10func 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 }
11func iabs(v: i64) -> i64 { if v<0 { return 0-v } return v }
12
13const S: i64 = 1048576 // Q20 -- the BACKWARD needs more than Q16 for small gradients through the softmax Jacobian
14const T: i64 = 3
15const DM: i64 = 4
16const SCALE: i64 = 524288 // 1/sqrt(4) = 0.5 in Q20
17
18// fixed-point exp, constants scaled to Q20 (log2e*S and quartic coeffs a1..a4 * S)
19func fp_exp_q16(xq: i64) -> i64 {
20 let y: i64=(xq*1512776)/S
21 var yi: i64=0
22 if y>=0 { yi=y/S } else { yi=0-(((0-y)+S-1)/S) }
23 let yf: i64=y-yi*S
24 var p: i64=10085
25 p=58197+(p*yf)/S; p=251882+(p*yf)/S; p=726817+(p*yf)/S; p=S+(p*yf)/S
26 if yi>=0 { if yi>=31 { return 2000000000 } return p*(1<<yi) }
27 let k: i64=0-yi; if k>=31 { return 0 }
28 return p/(1<<k)
29}
30
31// forward: fills Q,K,V,A,O; returns L_q32
32func attn_fwd(X: *i64, Wq: *i64, Wk: *i64, Wv: *i64, Q: *i64, K: *i64, V: *i64, A: *i64, O: *i64) -> i64 {
33 var t: i64=0
34 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
35 while k<DM { aq=aq+X[t*DM+k]*Wq[k*DM+i]; ak=ak+X[t*DM+k]*Wk[k*DM+i]; av=av+X[t*DM+k]*Wv[k*DM+i]; k=k+1 }
36 Q[t*DM+i]=aq/S; K[t*DM+i]=ak/S; V[t*DM+i]=av/S; i=i+1 } t=t+1 }
37 // scores (causal s<=t), softmax per row
38 t=0
39 while t<T {
40 var mx: i64=0-2000000000; var s: i64=0
41 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 }
42 var sum: i64=0; s=0; while s<=t { let e: i64=fp_exp_q16(A[t*T+s]-mx); A[t*T+s]=e; sum=sum+e; s=s+1 }
43 s=0; while s<=t { A[t*T+s]=(A[t*T+s]*S+sum/2)/sum; s=s+1 }
44 t=t+1
45 }
46 // O = A @ V (causal)
47 var Lq: i64=0; t=0
48 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 } let ov: i64=acc/S; O[t*DM+i]=ov; Lq=Lq+ov*ov; i=i+1 } t=t+1 }
49 return Lq
50}
51
52// backward: fills dWq,dWk,dWv (needs Q,K,V,A,O from a base-point forward)
53func attn_bwd(X: *i64, Q: *i64, K: *i64, V: *i64, A: *i64, O: *i64, dWq: *i64, dWk: *i64, dWv: *i64) -> i64 {
54 let dO: *i64=sys_mmap(T*DM*8) as *i64; let dV: *i64=sys_mmap(T*DM*8) as *i64
55 let dA: *i64=sys_mmap(T*T*8) as *i64; let dsc: *i64=sys_mmap(T*T*8) as *i64
56 let dQ: *i64=sys_mmap(T*DM*8) as *i64; let dK: *i64=sys_mmap(T*DM*8) as *i64
57 var t: i64=0; while t<T { var i: i64=0; while i<DM { dO[t*DM+i]=2*O[t*DM+i]; i=i+1 } t=t+1 }
58 // dV[s,i] = sum_{t>=s} A[t,s] dO[t,i]
59 var s: i64=0; while s<DM*T { dV[s]=0; s=s+1 }
60 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 }
61 // dA[t,s] = sum_i dO[t,i] V[s,i] (s<=t)
62 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 }
63 // dsc[t,s] = A[t,s]*(dA[t,s]-dot_t)/S ; dot_t = sum_s A[t,s]dA[t,s]/S (softmax Jacobian)
64 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 }
65 // dQ[t,i] = sum_{s<=t} (scale*dsc[t,s]) K[s,i] ; dK[s,i] = sum_{t>=s} (scale*dsc[t,s]) Q[t,i]
66 t=0; while t<T*DM { dQ[t]=0; dK[t]=0; t=t+1 }
67 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 }
68 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 }
69 // dW[k,i] = sum_t X[t,k] dZ[t,i]
70 var kk: i64=0
71 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+(X[t*DM+kk]*dQ[t*DM+i])/S; ak=ak+(X[t*DM+kk]*dK[t*DM+i])/S; av=av+(X[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 }
72 return 0
73}
74
75func gcheck(name: *u8, X: *i64, Wq: *i64, Wk: *i64, Wv: *i64, Wtarget: *i64, dW: *i64, Q: *i64, K: *i64, V: *i64, A: *i64, O: *i64) -> i64 {
76 let DELTA: i64=10486; let TOLP: i64=60 // delta ~0.01 at Q20; tol = 6% of the MATRIX's max |gradient| (standard gradcheck metric)
77 var maxabs: i64=1; var j: i64=0; while j<DM*DM { if iabs(dW[j])>maxabs { maxabs=iabs(dW[j]) } j=j+1 }
78 var npass: i64=0; var worst: i64=0; var i: i64=0
79 while i<DM*DM {
80 let save: i64=Wtarget[i]
81 Wtarget[i]=save+DELTA; let Lp: i64=attn_fwd(X,Wq,Wk,Wv,Q,K,V,A,O)
82 Wtarget[i]=save-DELTA; let Lm: i64=attn_fwd(X,Wq,Wk,Wv,Q,K,V,A,O)
83 Wtarget[i]=save; let dd: i64=attn_fwd(X,Wq,Wk,Wv,Q,K,V,A,O)
84 let num: i64=(Lp-Lm)/(2*DELTA); let rel: i64=(iabs(num-dW[i])*1000)/maxabs // error relative to max gradient
85 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) }
86 if rel>worst { worst=rel }
87 i=i+1
88 }
89 w(" " as *u8); w(name); w(": " as *u8); wn(npass); w("/" as *u8); wn(DM*DM); w(" (err<=6% of max|grad|=" as *u8); wn(maxabs); w("), worst=" as *u8); wn(worst); w("permil\n" as *u8)
90 return npass
91}
92
93func main() -> i64 {
94 w("=== nx_intfp_attention_gradcheck: Q16 single-head CAUSAL attention, full backward, gradcheck Wq/Wk/Wv -- no float ===\n\n" as *u8)
95 let X: *i64=sys_mmap(T*DM*8) as *i64
96 let Wq: *i64=sys_mmap(DM*DM*8) as *i64; let Wk: *i64=sys_mmap(DM*DM*8) as *i64; let Wv: *i64=sys_mmap(DM*DM*8) as *i64
97 let Q: *i64=sys_mmap(T*DM*8) as *i64; let K: *i64=sys_mmap(T*DM*8) as *i64; let V: *i64=sys_mmap(T*DM*8) as *i64
98 let A: *i64=sys_mmap(T*T*8) as *i64; let O: *i64=sys_mmap(T*DM*8) as *i64
99 let dWq: *i64=sys_mmap(DM*DM*8) as *i64; let dWk: *i64=sys_mmap(DM*DM*8) as *i64; let dWv: *i64=sys_mmap(DM*DM*8) as *i64
100 // small deterministic init
101 var i: i64=0; while i<T*DM { X[i]=((((i*5+2)%11)-5)*S)/12; i=i+1 }
102 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 }
103
104 let L0: i64=attn_fwd(X,Wq,Wk,Wv,Q,K,V,A,O)
105 attn_bwd(X,Q,K,V,A,O,dWq,dWk,dWv)
106 w(" forward L_q32=" as *u8); wn(L0); w(" (attn_scores 1/sqrt(4)=0.5; causal T=" as *u8); wn(T); w(" d=" as *u8); wn(DM); w(")\n" as *u8)
107 w(" gradcheck (only failing cells printed):\n" as *u8)
108 let pv: i64=gcheck("Wv" as *u8, X, Wq, Wk, Wv, Wv, dWv, Q,K,V,A,O) // easy path (O<-V<-Wv)
109 let pk: i64=gcheck("Wk" as *u8, X, Wq, Wk, Wv, Wk, dWk, Q,K,V,A,O) // thru softmax
110 let pq: i64=gcheck("Wq" as *u8, X, Wq, Wk, Wv, Wq, dWq, Q,K,V,A,O) // thru softmax
111 let tot: i64=pv+pk+pq; let want: i64=3*DM*DM
112 w("\n ATTENTION composition gradcheck: " as *u8); wn(tot); w("/" as *u8); wn(want); w(" cells correct\n" as *u8)
113 w("NX-INTFP-ATTENTION-GRADCHECK verdict=" as *u8)
114 if tot==want { w("GREEN " as *u8); wn(tot); w("/" as *u8); wn(want); w(" -- integer ops COMPOSE (softmax Jacobian + matmul-transpose backprops); the transformer tape is assemblable\n" as *u8) }
115 else { w("RED " as *u8); wn(tot); w("/" as *u8); wn(want); w(" -- a composition Q-scaling/index bug (see failing cells)\n" as *u8) }
116 return 0
117}