code wiki / _hdl_build / nx_intfp_rmsnorm_gradcheck_gate.nx
nx_intfp_rmsnorm_gradcheck_gate.nx source
↩ module page · 92 lines · 5338 B
1// nx_intfp_rmsnorm_gradcheck_gate.nx -- LAST transcendental for integer transformer training: RMSNorm needs
2// 1/sqrt, exactly where the software-float tape uses f32_sqrtx. Here: integer isqrt (exact bit-by-bit) + RMSNorm
3// forward + the full coupled Jacobian backward, gradchecked (integer finite-diff), NO float.
4// fwd: ms=mean(x^2)+eps ; rms=sqrt(ms) ; y_i = (x_i/rms)*gamma_i
5// bwd: dL/dx_j = gamma_j*g_j/rms - x_j*c/(D*rms^3), c=sum_i g_i*gamma_i*x_i, g_i=dL/dy_i
6// Key Q16 trick: rms_q16 = isqrt(ms_q32) exactly (isqrt(ms*S^2)=sqrt(ms)*S=rms*S). Uses inv_rms to avoid rms^3
7// overflow. If this gradchecks, the last hard op is proven and the whole integer tape is assemblable. license_tier: ORIGINAL
8import "nx_syscalls.nx"
9
10func w(s: *u8) -> i64 { var n: i64=0; while s[n]!=(0 as u8){n=n+1} sys_write(1,s,n); return 0 }
11func 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 }
12func iabs(v: i64) -> i64 { if v<0 { return 0-v } return v }
13
14const S: i64 = 65536
15const D: i64 = 4
16
17// exact integer floor-sqrt (bit-by-bit), O(log n)
18func isqrt(n: i64) -> i64 {
19 if n<=0 { return 0 }
20 var bit: i64=1; while bit*4<=n { bit=bit*4 }
21 var res: i64=0; var num: i64=n
22 while bit!=0 {
23 if num>=res+bit { num=num-(res+bit); res=(res/2)+bit } else { res=res/2 }
24 bit=bit/4
25 }
26 return res
27}
28
29// RMSNorm forward; fills y (Q16); returns rms_q16 (via out slot rmsout[0]) and L_q32=sum(y^2) as return
30func rms_fwd(x: *i64, gamma: *i64, y: *i64, rmsout: *i64) -> i64 {
31 var ms: i64=0; var i: i64=0
32 while i<D { ms=ms + x[i]*x[i]; i=i+1 } // sum x_iq^2 = Q32 of sum(x^2)
33 ms=ms/D + 42950 // mean + eps(1e-5 in Q32 = 42950)
34 let rms: i64=isqrt(ms) // = sqrt(ms)*S = rms in Q16
35 rmsout[0]=rms
36 let inv: i64=(S*S)/rms // 1/rms in Q16
37 var Lq: i64=0; i=0
38 while i<D { let nrm: i64=(x[i]*inv)/S; let yv: i64=(nrm*gamma[i])/S; y[i]=yv; Lq=Lq+yv*yv; i=i+1 }
39 return Lq
40}
41
42func main() -> i64 {
43 w("=== nx_intfp_rmsnorm_gradcheck: Q16 RMSNorm + isqrt + coupled Jacobian backward -- no float ===\n\n" as *u8)
44 // (0) verify isqrt before building on it
45 w(" [isqrt] isqrt(16)=" as *u8); wn(isqrt(16)); w(" (4) isqrt(1000000)=" as *u8); wn(isqrt(1000000)); w(" (1000) isqrt(2*S*S)=" as *u8); wn(isqrt(2*S*S)); w(" (~92682=1.414*S)\n\n" as *u8)
46
47 let x: *i64=sys_mmap(D*8) as *i64; let gamma: *i64=sys_mmap(D*8) as *i64
48 let y: *i64=sys_mmap(D*8) as *i64; let rmsout: *i64=sys_mmap(8) as *i64
49 x[0]=(S)/2; x[1]=(0-S*3)/10; x[2]=(S*4)/5; x[3]=(0-S*3)/5 // 0.5,-0.3,0.8,-0.6
50 gamma[0]=S; gamma[1]=(S*11)/10; gamma[2]=(S*9)/10; gamma[3]=(S*105)/100 // 1.0,1.1,0.9,1.05
51
52 let L0: i64=rms_fwd(x, gamma, y, rmsout); let rms: i64=rmsout[0]
53 w(" rms_q16=" as *u8); wn(rms); w(" y_q=[" as *u8); var i: i64=0; while i<D { wn(y[i]); if i<D-1 { w("," as *u8) } i=i+1 } w("] L_q32=" as *u8); wn(L0); w("\n\n" as *u8)
54
55 // analytic backward
56 let g: *i64=sys_mmap(D*8) as *i64; i=0; while i<D { g[i]=2*y[i]; i=i+1 }
57 let inv: i64=(S*S)/rms // 1/rms Q16
58 let invr3: i64=(((inv*inv)/S)*inv)/S // inv_rms^3 Q16
59 var c: i64=0; i=0; while i<D { let gg: i64=(g[i]*gamma[i])/S; let ggx: i64=(gg*x[i])/S; c=c+ggx; i=i+1 } // c_q16
60 let ana: *i64=sys_mmap(D*8) as *i64
61 i=0
62 while i<D {
63 let term1: i64=(((gamma[i]*g[i])/S)*inv)/S // gamma_j g_j / rms
64 let t: i64=(x[i]*c)/S // x_j c
65 let term2: i64=(((t*invr3)/S))/D // x_j c inv_rms^3 / D
66 ana[i]=term1-term2
67 i=i+1
68 }
69
70 let DELTA: i64=655; let TOLP: i64=80
71 var npass: i64=0; var worst: i64=0
72 w(" j analytic_q numeric_q rel(permille)\n" as *u8)
73 w(" ------------------------------------------------\n" as *u8)
74 i=0
75 while i<D {
76 let save: i64=x[i]
77 x[i]=save+DELTA; let Lp: i64=rms_fwd(x, gamma, y, rmsout)
78 x[i]=save-DELTA; let Lm: i64=rms_fwd(x, gamma, y, rmsout)
79 x[i]=save; let dd: i64=rms_fwd(x, gamma, y, rmsout)
80 let num: i64=(Lp-Lm)/(2*DELTA)
81 let rel: i64=(iabs(num-ana[i])*1000)/(iabs(ana[i])+100)
82 w(" " as *u8); wn(i); w(" " as *u8); wn(ana[i]); w(" " as *u8); wn(num); w(" " as *u8); wn(rel)
83 if rel<=TOLP { npass=npass+1; w(" ok\n" as *u8) } else { w(" FAIL\n" as *u8) }
84 if rel>worst { worst=rel }
85 i=i+1
86 }
87 w("\n RMSNorm gradcheck: " as *u8); wn(npass); w("/" as *u8); wn(D); w(" within " as *u8); wn(TOLP); w(" permille; worst=" as *u8); wn(worst); w("\n" as *u8)
88 w("NX-INTFP-RMSNORM-GRADCHECK verdict=" as *u8)
89 if npass==D { w("GREEN passes=" as *u8); wn(npass); w("/" as *u8); wn(D); w(" -- integer RMSNorm+isqrt+coupled-Jacobian proven. ALL hard transcendentals (exp, rsqrt) now integer.\n" as *u8) }
90 else { w("RED passes=" as *u8); wn(npass); w("/" as *u8); wn(D); w(" -- isqrt or Jacobian Q-scaling bug\n" as *u8) }
91 return 0
92}