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}