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}