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}