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}