code wiki / _hdl_build / nx_intfp_mla_gradcheck_gate.nx
nx_intfp_mla_gradcheck_gate.nx source
↩ module page · 98 lines · 9318 B
1// nx_intfp_mla_gradcheck_gate.nx -- 2026 FRONTIER: MLA (Multi-head Latent Attention, DeepSeek's marquee KV-
2// compression) in Q20 INTEGER, gradchecked, NO float. Instead of storing full K,V, x is DOWN-projected to a small
3// latent c (dim LC<DM), and K,V are UP-projected from c -- so the KV cache is just c (LC-dim), a big compression.
4// Full causal attention on the reconstructed K,V, full backward through the bottleneck. Gradcheck the DOWN-proj
5// Wdkv (the compression novelty -- its gradient flows through BOTH the K and V up-projections and all of attention)
6// plus Wuk, Wuv, Wq. 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 = 16777216 // Q24 -- Q/K gradients flow thru softmax and are tiny; need the extra bits (as attention gate)
14const T: i64 = 3
15const DM: i64 = 4
16const LC: i64 = 2 // latent (compressed) dim < DM -- the KV-cache size
17const SCALE: i64 = 8388608 // 1/sqrt(4)=0.5
18
19// fp_exp constants scaled to Q24 (log2e*S + quartic a1..a4*S)
20func 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) }
21
22// slots: 0x 1Wdkv 2Wuk 3Wuv 4Wq | 5c 6K 7V 8Q 9A 10O | 11dWdkv 12dWuk 13dWuv 14dWq
23func mla_fwd(P: *i64) -> i64 {
24 let x: *i64=P[0] as *i64; let Wdkv: *i64=P[1] as *i64; let Wuk: *i64=P[2] as *i64; let Wuv: *i64=P[3] as *i64; let Wq: *i64=P[4] as *i64
25 let c: *i64=P[5] as *i64; let K: *i64=P[6] as *i64; let V: *i64=P[7] as *i64; let Q: *i64=P[8] as *i64; let A: *i64=P[9] as *i64; let O: *i64=P[10] as *i64
26 // down-project to latent c
27 var t: i64=0; while t<T { var l: i64=0; while l<LC { var acc: i64=0; var k: i64=0; while k<DM { acc=acc+x[t*DM+k]*Wdkv[k*LC+l]; k=k+1 } c[t*LC+l]=acc/S; l=l+1 } t=t+1 }
28 // up-project K,V from c ; Q from x
29 t=0; while t<T { var i: i64=0; while i<DM { var ak: i64=0; var av: i64=0; var l: i64=0; while l<LC { ak=ak+c[t*LC+l]*Wuk[l*DM+i]; av=av+c[t*LC+l]*Wuv[l*DM+i]; l=l+1 } K[t*DM+i]=ak/S; V[t*DM+i]=av/S; var aq: i64=0; var k: i64=0; while k<DM { aq=aq+x[t*DM+k]*Wq[k*DM+i]; k=k+1 } Q[t*DM+i]=aq/S; i=i+1 } t=t+1 }
30 // causal scaled-dot attention on reconstructed K,V
31 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+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 } 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 }
32 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 }
33 return Lq
34}
35func mla_bwd(P: *i64) -> i64 {
36 let x: *i64=P[0] as *i64; let Wuk: *i64=P[2] as *i64; let Wuv: *i64=P[3] as *i64; let Wq: *i64=P[4] as *i64
37 let c: *i64=P[5] as *i64; let K: *i64=P[6] as *i64; let V: *i64=P[7] as *i64; let Q: *i64=P[8] as *i64; let A: *i64=P[9] as *i64; let O: *i64=P[10] as *i64
38 let dWdkv: *i64=P[11] as *i64; let dWuk: *i64=P[12] as *i64; let dWuv: *i64=P[13] as *i64; let dWq: *i64=P[14] as *i64
39 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
40 let dQ: *i64=sys_mmap(T*DM*8) as *i64; let dK: *i64=sys_mmap(T*DM*8) as *i64; let dc: *i64=sys_mmap(T*LC*8) as *i64
41 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 }
42 var s: i64=0; while s<DM*T { dV[s]=0; s=s+1 }
43 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 }
44 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 }
45 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 }
46 t=0; while t<T*DM { dQ[t]=0; dK[t]=0; t=t+1 }
47 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 }
48 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 }
49 // dWq[k,i]=sum_t x[t,k]dQ[t,i]
50 var kk: i64=0; while kk<DM { var i: i64=0; while i<DM { var acc: i64=0; t=0; while t<T { acc=acc+(x[t*DM+kk]*dQ[t*DM+i])/S; t=t+1 } dWq[kk*DM+i]=acc; i=i+1 } kk=kk+1 }
51 // through the KV bottleneck: dWuk[l,i]=sum_t c[t,l]dK[t,i] ; dWuv[l,i]=sum_t c[t,l]dV[t,i] ; dc = dK@Wuk^T + dV@Wuv^T
52 var l: i64=0; while l<LC { var i: i64=0; while i<DM { var auk: i64=0; var auv: i64=0; t=0; while t<T { auk=auk+(c[t*LC+l]*dK[t*DM+i])/S; auv=auv+(c[t*LC+l]*dV[t*DM+i])/S; t=t+1 } dWuk[l*DM+i]=auk; dWuv[l*DM+i]=auv; i=i+1 } l=l+1 }
53 t=0; while t<T { l=0; while l<LC { var acc: i64=0; var i: i64=0; while i<DM { acc=acc+(dK[t*DM+i]*Wuk[l*DM+i])/S+(dV[t*DM+i]*Wuv[l*DM+i])/S; i=i+1 } dc[t*LC+l]=acc; l=l+1 } t=t+1 }
54 // dWdkv[k,l]=sum_t x[t,k]dc[t,l]
55 kk=0; while kk<DM { l=0; while l<LC { var acc: i64=0; t=0; while t<T { acc=acc+(x[t*DM+kk]*dc[t*LC+l])/S; t=t+1 } dWdkv[kk*LC+l]=acc; l=l+1 } kk=kk+1 }
56 return 0
57}
58func gcheck(name: *u8, P: *i64, Wt: *i64, dW: *i64, ncell: i64) -> i64 {
59 let DELTA: i64=33552; let TOLP: i64=70 // ~0.002 at Q24
60 var maxabs: i64=1; var q: i64=0; while q<ncell { if iabs(dW[q])>maxabs { maxabs=iabs(dW[q]) } q=q+1 }
61 var npass: i64=0; var worst: i64=0; var i: i64=0
62 while i<ncell {
63 let save: i64=Wt[i]
64 Wt[i]=save+DELTA; let Lp: i64=mla_fwd(P)
65 Wt[i]=save-DELTA; let Lm: i64=mla_fwd(P)
66 Wt[i]=save; let dd: i64=mla_fwd(P)
67 let num: i64=(Lp-Lm)/(2*DELTA); let rel: i64=(iabs(num-dW[i])*1000)/maxabs
68 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) }
69 if rel>worst { worst=rel }
70 i=i+1
71 }
72 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)
73 return npass
74}
75func main() -> i64 {
76 w("=== nx_intfp_mla_gradcheck: 2026-FRONTIER Multi-head Latent Attention (KV-compression, DeepSeek) Q20 INTEGER ===\n\n" as *u8)
77 let P: *i64=sys_mmap(15*8) as *i64
78 P[0]=sys_mmap(T*DM*8); P[1]=sys_mmap(DM*LC*8); P[2]=sys_mmap(LC*DM*8); P[3]=sys_mmap(LC*DM*8); P[4]=sys_mmap(DM*DM*8)
79 P[5]=sys_mmap(T*LC*8); 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*T*8); P[10]=sys_mmap(T*DM*8)
80 P[11]=sys_mmap(DM*LC*8); P[12]=sys_mmap(LC*DM*8); P[13]=sys_mmap(LC*DM*8); P[14]=sys_mmap(DM*DM*8)
81 let x: *i64=P[0] as *i64; let Wdkv: *i64=P[1] as *i64; let Wuk: *i64=P[2] as *i64; let Wuv: *i64=P[3] as *i64; let Wq: *i64=P[4] as *i64
82 var i: i64=0; while i<T*DM { x[i]=((((i*5+2)%11)-5)*S)/10; i=i+1 }
83 i=0; while i<DM*LC { Wdkv[i]=((((i*7+1)%13)-6)*S)/14; i=i+1 }
84 i=0; while i<LC*DM { Wuk[i]=((((i*3+5)%13)-6)*S)/14; Wuv[i]=((((i*11+2)%13)-6)*S)/14; i=i+1 }
85 i=0; while i<DM*DM { Wq[i]=((((i*5+3)%13)-6)*S)/14; i=i+1 }
86 let L0: i64=mla_fwd(P); mla_bwd(P)
87 w(" KV compression: DM=" as *u8); wn(DM); w(" -> latent LC=" as *u8); wn(LC); w(" (KV cache " as *u8); wn((LC*100)/(2*DM)); w("% of full K+V) forward L_q32=" as *u8); wn(L0); w("\n gradcheck:\n" as *u8)
88 let pd: i64=gcheck("Wdkv(down)" as *u8, P, Wdkv, P[11] as *i64, DM*LC) // THE compression novelty
89 let puk: i64=gcheck("Wuk (upK) " as *u8, P, Wuk, P[12] as *i64, LC*DM)
90 let puv: i64=gcheck("Wuv (upV) " as *u8, P, Wuv, P[13] as *i64, LC*DM)
91 let pq: i64=gcheck("Wq (qry) " as *u8, P, Wq, P[14] as *i64, DM*DM)
92 let tot: i64=pd+puk+puv+pq; let want: i64=DM*LC+LC*DM+LC*DM+DM*DM
93 w("\n MLA gradcheck: " as *u8); wn(tot); w("/" as *u8); wn(want); w(" cells correct\n" as *u8)
94 w("NX-INTFP-MLA verdict=" as *u8)
95 if tot==want { w("GREEN " as *u8); wn(tot); w("/" as *u8); wn(want); w(" -- 2026-FRONTIER integer MLA (KV-compression thru latent) PROVEN, no float. DeepSeek's marquee attention on the sovereign stack.\n" as *u8) }
96 else { w("RED " as *u8); wn(tot); w("/" as *u8); wn(want); w(" -- bottleneck Q-scaling bug\n" as *u8) }
97 return 0
98}