code wiki / _hdl_build / nx_intfp_qknorm_gradcheck_gate.nx

nx_intfp_qknorm_gradcheck_gate.nx source

↩ module page · 101 lines · 9960 B

1// nx_intfp_qknorm_gradcheck_gate.nx -- 2026 FRONTIER: QK-NORM (Qwen3.5) in Q24 INTEGER, gradchecked, no float. 2// Applies RMSNorm (with learned gamma) to Q and K per-token over the head dim BEFORE the attention dot product -- 3// stabilizes attention logits at scale (kills the exploding-logit failure mode). Composes proven RMSNorm + causal 4// attention; backward = attention-bwd -> dQn/dKn -> RMSNorm-bwd -> dWq/dWk + dgq/dgk. Gradcheck Wq (thru Q-norm), 5// gq (Q-norm gamma, the novelty), Wv. 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 = 16777216 // Q24 13const T: i64 = 3 14const DM: i64 = 4 15const SCALE: i64 = 8388608 16const EPS: i64 = 16777216 // eps in Q48 (tiny) 17 18func isqrt(n: i64) -> i64 { if n<=0 { return 0 } var bit: i64=1; while bit*4<=n { bit=bit*4 } var res: i64=0; var num: i64=n; while bit!=0 { if num>=res+bit { num=num-(res+bit); res=(res/2)+bit } else { res=res/2 } bit=bit/4 } return res } 19func fp_exp(xq: i64) -> i64 { let y: i64=(xq*24204406)/S; var yi: i64=0; if y>=0 { yi=y/S } else { yi=0-(((0-y)+S-1)/S) } let yf: i64=y-yi*S; var p: i64=161380; p=931144+(p*yf)/S; p=4030770+(p*yf)/S; p=11632166+(p*yf)/S; p=S+(p*yf)/S; if yi>=0 { if yi>=31 { return 2000000000 } return p*(1<<yi) } let k: i64=0-yi; if k>=31 { return 0 } return p/(1<<k) } 20 21// slots: 0x 1Wq 2Wk 3Wv 4gq 5gk | 6Q 7K 8V 9Qn 10Kn 11rq 12rk 13A 14O | 15dWq 16dWk 17dWv 18dgq 19dgk 22func qkn_fwd(P: *i64) -> i64 { 23 let x: *i64=P[0] as *i64; let Wq: *i64=P[1] as *i64; let Wk: *i64=P[2] as *i64; let Wv: *i64=P[3] as *i64; let gq: *i64=P[4] as *i64; let gk: *i64=P[5] as *i64 24 let Q: *i64=P[6] as *i64; let K: *i64=P[7] as *i64; let V: *i64=P[8] as *i64; let Qn: *i64=P[9] as *i64; let Kn: *i64=P[10] as *i64; let rq: *i64=P[11] as *i64; let rk: *i64=P[12] as *i64; let A: *i64=P[13] as *i64; let O: *i64=P[14] as *i64 25 var t: i64=0; 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; 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 } Q[t*DM+i]=aq/S; K[t*DM+i]=ak/S; V[t*DM+i]=av/S; i=i+1 } t=t+1 } 26 // QK-norm: RMSNorm over head dim with gamma 27 t=0; while t<T { var mq: i64=0; var mk: i64=0; var i: i64=0; while i<DM { mq=mq+Q[t*DM+i]*Q[t*DM+i]; mk=mk+K[t*DM+i]*K[t*DM+i]; i=i+1 } mq=mq/DM+EPS; mk=mk/DM+EPS; var r1: i64=isqrt(mq); if r1<1 { r1=1 } var r2: i64=isqrt(mk); if r2<1 { r2=1 } rq[t]=r1; rk[t]=r2; let iq: i64=(S*S)/r1; let ik: i64=(S*S)/r2; i=0; while i<DM { Qn[t*DM+i]=(((Q[t*DM+i]*iq)/S)*gq[i])/S; Kn[t*DM+i]=(((K[t*DM+i]*ik)/S)*gk[i])/S; i=i+1 } t=t+1 } 28 // causal attention on normalized Q,K 29 t=0; while t<T { var mx: i64=0-2000000000; var s: i64=0; while s<=t { var dot: i64=0; var i: i64=0; while i<DM { dot=dot+Qn[t*DM+i]*Kn[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 } 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 } s=0; while s<=t { A[t*T+s]=(A[t*T+s]*S+sum/2)/sum; s=s+1 } t=t+1 } 30 var Lq: i64=0; t=0; 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 } 31 return Lq 32} 33func qkn_bwd(P: *i64) -> i64 { 34 let x: *i64=P[0] as *i64; let gq: *i64=P[4] as *i64; let gk: *i64=P[5] as *i64 35 let Q: *i64=P[6] as *i64; let K: *i64=P[7] as *i64; let V: *i64=P[8] as *i64; let Qn: *i64=P[9] as *i64; let Kn: *i64=P[10] as *i64; let rq: *i64=P[11] as *i64; let rk: *i64=P[12] as *i64; let A: *i64=P[13] as *i64; let O: *i64=P[14] as *i64 36 let dWq: *i64=P[15] as *i64; let dWk: *i64=P[16] as *i64; let dWv: *i64=P[17] as *i64; let dgq: *i64=P[18] as *i64; let dgk: *i64=P[19] as *i64 37 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 38 let dQn: *i64=sys_mmap(T*DM*8) as *i64; let dKn: *i64=sys_mmap(T*DM*8) as *i64; let dQ: *i64=sys_mmap(T*DM*8) as *i64; let dK: *i64=sys_mmap(T*DM*8) as *i64 39 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 } 40 var s: i64=0; while s<DM*T { dV[s]=0; s=s+1 } 41 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 } 42 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 } 43 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 } 44 // dQn[t,i]=sum_{s<=t}(scale*dsc)Kn[s,i] ; dKn[s,i]=sum_{t>=s}(scale*dsc)Qn[t,i] 45 t=0; while t<T*DM { dQn[t]=0; dKn[t]=0; t=t+1 } 46 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*Kn[s*DM+i])/S; s=s+1 } dQn[t*DM+i]=acc; i=i+1 } t=t+1 } 47 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*Qn[t*DM+i])/S; t=t+1 } dKn[s*DM+i]=acc; i=i+1 } s=s+1 } 48 // RMSNorm backward on Q (with gamma gq) and K (gk): dQn -> dQ + dgq 49 var i2: i64=0; while i2<DM { dgq[i2]=0; dgk[i2]=0; i2=i2+1 } 50 t=0 51 while t<T { 52 let r1: i64=rq[t]; let iq: i64=(S*S)/r1; let iq3: i64=(((iq*iq)/S)*iq)/S 53 let r2: i64=rk[t]; let ik: i64=(S*S)/r2; let ik3: i64=(((ik*ik)/S)*ik)/S 54 var cq: i64=0; var ck: i64=0; var i: i64=0 55 while i<DM { let nq: i64=(Q[t*DM+i]*iq)/S; dgq[i]=dgq[i]+(dQn[t*DM+i]*nq)/S; let dnq: i64=(dQn[t*DM+i]*gq[i])/S; cq=cq+(dnq*Q[t*DM+i])/S; let nk: i64=(K[t*DM+i]*ik)/S; dgk[i]=dgk[i]+(dKn[t*DM+i]*nk)/S; let dnk: i64=(dKn[t*DM+i]*gk[i])/S; ck=ck+(dnk*K[t*DM+i])/S; i=i+1 } 56 i=0; while i<DM { let dnq: i64=(dQn[t*DM+i]*gq[i])/S; let tq1: i64=(dnq*iq)/S; let ttq: i64=(Q[t*DM+i]*cq)/S; let tq2: i64=(((ttq*iq3)/S))/DM; dQ[t*DM+i]=tq1-tq2; let dnk: i64=(dKn[t*DM+i]*gk[i])/S; let tk1: i64=(dnk*ik)/S; let ttk: i64=(K[t*DM+i]*ck)/S; let tk2: i64=(((ttk*ik3)/S))/DM; dK[t*DM+i]=tk1-tk2; i=i+1 } 57 t=t+1 58 } 59 // dWq=sum_t x*dQ ; dWk=sum_t x*dK ; dWv=sum_t x*dV 60 var kk: i64=0; 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 } 61 return 0 62} 63func gcheck(name: *u8, P: *i64, Wt: *i64, dW: *i64, ncell: i64) -> i64 { 64 let DELTA: i64=33552; let TOLP: i64=70 65 var maxabs: i64=1; var q: i64=0; while q<ncell { if iabs(dW[q])>maxabs { maxabs=iabs(dW[q]) } q=q+1 } 66 var npass: i64=0; var worst: i64=0; var i: i64=0 67 while i<ncell { 68 let save: i64=Wt[i] 69 Wt[i]=save+DELTA; let Lp: i64=qkn_fwd(P) 70 Wt[i]=save-DELTA; let Lm: i64=qkn_fwd(P) 71 Wt[i]=save; let dd: i64=qkn_fwd(P) 72 let num: i64=(Lp-Lm)/(2*DELTA); let rel: i64=(iabs(num-dW[i])*1000)/maxabs 73 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) } 74 if rel>worst { worst=rel } 75 i=i+1 76 } 77 w(" " as *u8); w(name); w(": " as *u8); wn(npass); w("/" as *u8); wn(ncell); w(" (worst=" as *u8); wn(worst); w("permil)\n" as *u8) 78 return npass 79} 80func main() -> i64 { 81 w("=== nx_intfp_qknorm_gradcheck: 2026-FRONTIER QK-Norm (Qwen3.5) Q24 INTEGER, no float ===\n\n" as *u8) 82 let P: *i64=sys_mmap(20*8) as *i64 83 P[0]=sys_mmap(T*DM*8); P[1]=sys_mmap(DM*DM*8); P[2]=sys_mmap(DM*DM*8); P[3]=sys_mmap(DM*DM*8); P[4]=sys_mmap(DM*8); P[5]=sys_mmap(DM*8) 84 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*DM*8); P[11]=sys_mmap(T*8); P[12]=sys_mmap(T*8); P[13]=sys_mmap(T*T*8); P[14]=sys_mmap(T*DM*8) 85 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(DM*8); P[19]=sys_mmap(DM*8) 86 let x: *i64=P[0] as *i64; let Wq: *i64=P[1] as *i64; let Wk: *i64=P[2] as *i64; let Wv: *i64=P[3] as *i64; let gq: *i64=P[4] as *i64; let gk: *i64=P[5] as *i64 87 var i: i64=0; while i<T*DM { x[i]=((((i*5+2)%11)-5)*S)/10; i=i+1 } 88 i=0; while i<DM*DM { Wq[i]=((((i*7+1)%13)-6)*S)/14; Wk[i]=((((i*3+5)%13)-6)*S)/14; Wv[i]=((((i*11+2)%13)-6)*S)/14; i=i+1 } 89 i=0; while i<DM { gq[i]=S; gk[i]=S; i=i+1 } 90 let L0: i64=qkn_fwd(P); qkn_bwd(P) 91 w(" forward L_q32=" as *u8); wn(L0); w(" (RMSNorm on Q,K over head dim before scores)\n gradcheck:\n" as *u8) 92 let pq: i64=gcheck("Wq " as *u8, P, Wq, P[15] as *i64, DM*DM) 93 let pgq: i64=gcheck("gq " as *u8, P, gq, P[18] as *i64, DM) // Q-norm gamma (the novelty) 94 let pv: i64=gcheck("Wv " as *u8, P, Wv, P[17] as *i64, DM*DM) 95 let tot: i64=pq+pgq+pv; let want: i64=DM*DM+DM+DM*DM 96 w("\n QK-Norm gradcheck: " as *u8); wn(tot); w("/" as *u8); wn(want); w(" cells correct\n" as *u8) 97 w("NX-INTFP-QKNORM verdict=" as *u8) 98 if tot==want { w("GREEN " as *u8); wn(tot); w("/" as *u8); wn(want); w(" -- 2026-FRONTIER integer QK-Norm PROVEN (RMSNorm on Q,K + gamma, thru attention), no float.\n" as *u8) } 99 else { w("RED " as *u8); wn(tot); w("/" as *u8); wn(want); w(" -- QK-norm Q-scaling bug\n" as *u8) } 100 return 0 101}