code wiki / _hdl_build / nx_f32_rmsnorm_rope_gate.nx

nx_f32_rmsnorm_rope_gate.nx source

↩ module page · 133 lines · 10370 B

1import "nx_gate_gn.nx" 2import "nx_gate_base.nx" 3// nx_f32_rmsnorm_rope_gate.nx -- RUNG 3 (Llama-3.2 architecture): RMSNorm + RoPE backward, gradient-checked. Llama 4// uses RMSNorm (not layernorm: NO mean subtraction, scale-invariant) and RoPE (rotary position embedding: a 2D 5// rotation per dim-pair, so its backward is the INVERSE rotation -- transpose). Adds f32_sin/f32_cos (Taylor). Same 6// finite-difference proof discipline as R1/R2. Targets the lightweight edge arch we're exceeding Llama on (sovereign, 7// deterministic, hardware-universal). 8// T0 f32_sin/cos: sin(0.5)=0.4794, cos(0.5)=0.8776 (Taylor). 9// T1 RMSNORM FORWARD: xhat = x/rms(x); rms([1,2,3])=2.16 -> xhat=[0.463,0.926,1.389]. 10// T2 RMSNORM dL/dx GRADCHECK: dx_k=(1/r)(dxhat_k - xhat_k*mean(dxhat*xhat)) == central finite-difference (no mean term, unlike layernorm). 11// T3 SCALE-INVARIANCE: sum_i dL/dx_i * x_i = 0 (scaling x by c leaves RMSNorm unchanged -> gradient orthogonal to x). 12// T4 RoPE BACKWARD GRADCHECK: rotation backward (inverse rotation) == finite-difference. 13// T5 NEGATIVE CONTROL + Llama-arch nonlinearities covered (with softmax R2a) -> next: GQA wiring + SwiGLU. 14// license_tier: ORIGINAL 15import "nx_f32_hw.nx" 16import "nx_syscalls.nx" 17 18func grow(name: *u8, ok: i64) -> i64 { if ok==1 { gw(" PASS " as *u8) } else { gw(" FAIL " as *u8) } gw(name); gw(" 19" as *u8); return ok } 20func gm(x: i64) -> i64 { return gn(f32_int(f32_mul(x, f32_of(1000)))) } 21func f32_abs(x: i64) -> i64 { return x & 0x7FFFFFFF } 22func f32_le(x: i64, y: i64) -> i64 { let d: i64=f32_sub(x,y) & 0xFFFFFFFF; if ((d>>31)&1)==1 { return 1 } if (d & 0x7FFFFFFF)==0 { return 1 } return 0 } 23func f32_sqrt(x: i64) -> i64 { if (x & 0x7FFFFFFF)==0 { return f32_of(0) } var y: i64=x; var i: i64=0; while i<14 { y=f32_div(f32_add(y, f32_div(x,y)), f32_of(2)); i=i+1 } return y } 24// f32_sin/cos via Taylor (small-angle; PRODUCTION RoPE needs range-reduction mod 2pi for large position*theta -- noted). 25func f32_sin(x: i64) -> i64 { let x2: i64=f32_mul(x,x); var term: i64=x; var sum: i64=x; var k: i64=1; while k<=8 { term=f32_neg(f32_div(f32_mul(term,x2), f32_of((2*k)*(2*k+1)))); sum=f32_add(sum,term); k=k+1 } return sum } 26func f32_cos(x: i64) -> i64 { let x2: i64=f32_mul(x,x); var term: i64=f32_of(1); var sum: i64=f32_of(1); var k: i64=1; while k<=8 { term=f32_neg(f32_div(f32_mul(term,x2), f32_of((2*k-1)*(2*k)))); sum=f32_add(sum,term); k=k+1 } return sum } 27 28// RMSNorm forward: r = sqrt(mean(x^2)+eps); xhat = x/r; y = gamma*xhat. (no mean subtraction) 29func rmsnorm_fwd(x: *i64, gamma: *i64, n: i64, eps: i64, y: *i64, xhat: *i64, rout: *i64) -> i64 { 30 var ss: i64=f32_of(0); var i: i64=0; while i<n { ss=f32_add(ss, f32_mul(x[i],x[i])); i=i+1 } 31 let r: i64=f32_sqrt(f32_add(f32_div(ss, f32_of(n)), eps)); rout[0]=r 32 i=0; while i<n { xhat[i]=f32_div(x[i], r); y[i]=f32_mul(gamma[i], xhat[i]); i=i+1 } 33 return 0 34} 35// RMSNorm backward: dx_k = (1/r)*(dxhat_k - xhat_k*mean(dxhat*xhat)), dxhat=g*gamma. 36func rmsnorm_bwd(g: *i64, gamma: *i64, xhat: *i64, r: i64, n: i64, dx: *i64) -> i64 { 37 let dxhat: *i64=sys_mmap(128) as *i64; var s2: i64=f32_of(0); var i: i64=0 38 while i<n { dxhat[i]=f32_mul(g[i],gamma[i]); s2=f32_add(s2, f32_mul(dxhat[i],xhat[i])); i=i+1 } 39 let mdxx: i64=f32_div(s2, f32_of(n)); let invr: i64=f32_div(f32_of(1), r) 40 i=0; while i<n { dx[i]=f32_mul(invr, f32_sub(dxhat[i], f32_mul(xhat[i], mdxx))); i=i+1 } 41 return 0 42} 43func rms_L(x: *i64, gamma: *i64, w: *i64, n: i64, eps: i64) -> i64 { 44 let y: *i64=sys_mmap(128) as *i64; let xh: *i64=sys_mmap(128) as *i64; let rr: *i64=sys_mmap(16) as *i64 45 rmsnorm_fwd(x, gamma, n, eps, y, xh, rr) 46 var L: i64=f32_of(0); var i: i64=0; while i<n { L=f32_add(L, f32_mul(w[i],y[i])); i=i+1 } 47 return L 48} 49// RoPE one pair: y0=x0*cos-x1*sin, y1=x0*sin+x1*cos. L=w0*y0+w1*y1 (forward, for the probe). 50func rope_L(x0: i64, x1: i64, a: i64, w0: i64, w1: i64) -> i64 { 51 let c: i64=f32_cos(a); let s: i64=f32_sin(a) 52 let y0: i64=f32_sub(f32_mul(x0,c), f32_mul(x1,s)); let y1: i64=f32_add(f32_mul(x0,s), f32_mul(x1,c)) 53 return f32_add(f32_mul(w0,y0), f32_mul(w1,y1)) 54} 55 56func main() -> i64 { 57 gw("=== nx_f32_rmsnorm_rope_gate: RUNG 3 -- RMSNorm + RoPE backward (gradchecked), the Llama-3.2 nonlinearities ===\n" as *u8) 58 var pass: i64=0; var total: i64=0 59 let eps: i64=f32_div(f32_of(1), f32_of(100000)); let h: i64=f32_div(f32_of(1), f32_of(100)); let tol: i64=f32_div(f32_of(3), f32_of(100)); let twoh: i64=f32_mul(f32_of(2), h) 60 61 // T0 sin/cos. 62 let sn: i64=f32_int(f32_mul(f32_sin(f32_div(f32_of(1),f32_of(2))),f32_of(1000))) 63 let cs: i64=f32_int(f32_mul(f32_cos(f32_div(f32_of(1),f32_of(2))),f32_of(1000))) 64 total=total+1; if sn>=477 { if sn<=481 { if cs>=875 { if cs<=879 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } 65 gw("T0 f32 sin/cos (milliunits): sin(0.5)=" as *u8); gn(sn); gw(" cos(0.5)=" as *u8); gn(cs); gw(" (expect 479/877)\n" as *u8) 66 67 // RMSNorm setup. 68 let x: *i64=sys_mmap(128) as *i64; x[0]=f32_of(1); x[1]=f32_of(2); x[2]=f32_of(3) 69 let gamma: *i64=sys_mmap(128) as *i64; gamma[0]=f32_of(1); gamma[1]=f32_of(1); gamma[2]=f32_of(1) 70 let w: *i64=sys_mmap(128) as *i64; w[0]=f32_of(1); w[1]=f32_of(0); w[2]=f32_of(0) 71 let N: i64=3 72 let y: *i64=sys_mmap(128) as *i64; let xhat: *i64=sys_mmap(128) as *i64; let rr: *i64=sys_mmap(16) as *i64 73 rmsnorm_fwd(x, gamma, N, eps, y, xhat, rr) 74 75 // T1 forward. 76 let xh0: i64=f32_int(f32_mul(xhat[0],f32_of(1000))) 77 total=total+1; if xh0>=460 { if xh0<=466 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } 78 gw("T1 RMSNORM FORWARD: xhat=x/rms = [" as *u8); gm(xhat[0]); gw("," as *u8); gm(xhat[1]); gw("," as *u8); gm(xhat[2]); gw("]m (expect 463,926,1389)\n" as *u8) 79 80 // T2 RMSNorm dL/dx gradcheck. 81 let dx: *i64=sys_mmap(128) as *i64; rmsnorm_bwd(w, gamma, xhat, rr[0], N, dx) 82 var gok: i64=1; var j: i64=0 83 while j<N { 84 let xp: *i64=sys_mmap(128) as *i64; let xm: *i64=sys_mmap(128) as *i64; var c: i64=0 85 while c<N { xp[c]=x[c]; xm[c]=x[c]; c=c+1 } 86 xp[j]=f32_add(x[j],h); xm[j]=f32_sub(x[j],h) 87 let fd: i64=f32_div(f32_sub(rms_L(xp,gamma,w,N,eps), rms_L(xm,gamma,w,N,eps)), twoh) 88 if f32_le(f32_abs(f32_sub(dx[j], fd)), tol)==0 { gok=0 } 89 gw(" d/dx" as *u8); gn(j); gw(": autograd=" as *u8); gm(dx[j]); gw("m finite-diff=" as *u8); gm(fd); gw("m\n" as *u8) 90 j=j+1 91 } 92 total=total+1; if gok==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } 93 gw("T2 RMSNORM dL/dx GRADCHECK: backward == finite-diff (no mean term, unlike layernorm)\n" as *u8) 94 95 // T3 scale-invariance: sum dL/dx_i * x_i ~= 0. 96 var dot: i64=f32_of(0); j=0; while j<N { dot=f32_add(dot, f32_mul(dx[j],x[j])); j=j+1 } 97 let doti: i64=f32_int(f32_mul(dot,f32_of(1000))) 98 total=total+1; if doti>=-3 { if doti<=3 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } 99 gw("T3 SCALE-INVARIANCE: sum(dL/dx * x)=" as *u8); gn(doti); gw("m ~= 0 (RMSNorm is scale-invariant -> gradient orthogonal to x)\n" as *u8) 100 101 // T4 RoPE backward gradcheck. pair x=[1,2], a=0.5, w=[1,0]. analytic: dx0=w0*cos+w1*sin, dx1=-w0*sin+w1*cos. 102 let rx0: i64=f32_of(1); let rx1: i64=f32_of(2); let aa: i64=f32_div(f32_of(1),f32_of(2)); let rw0: i64=f32_of(1); let rw1: i64=f32_of(0) 103 let cc: i64=f32_cos(aa); let ssn: i64=f32_sin(aa) 104 let adx0: i64=f32_add(f32_mul(rw0,cc), f32_mul(rw1,ssn)) 105 let adx1: i64=f32_add(f32_neg(f32_mul(rw0,ssn)), f32_mul(rw1,cc)) 106 let fdx0: i64=f32_div(f32_sub(rope_L(f32_add(rx0,h),rx1,aa,rw0,rw1), rope_L(f32_sub(rx0,h),rx1,aa,rw0,rw1)), twoh) 107 let fdx1: i64=f32_div(f32_sub(rope_L(rx0,f32_add(rx1,h),aa,rw0,rw1), rope_L(rx0,f32_sub(rx1,h),aa,rw0,rw1)), twoh) 108 var ropeok: i64=1 109 if f32_le(f32_abs(f32_sub(adx0,fdx0)),tol)==0 { ropeok=0 } 110 if f32_le(f32_abs(f32_sub(adx1,fdx1)),tol)==0 { ropeok=0 } 111 total=total+1; if ropeok==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } 112 gw("T4 RoPE BACKWARD GRADCHECK: dx0 ana=" as *u8); gm(adx0); gw("m fd=" as *u8); gm(fdx0); gw("m ; dx1 ana=" as *u8); gm(adx1); gw("m fd=" as *u8); gm(fdx1); gw("m (inverse rotation == finite-diff)\n" as *u8) 113 114 // T5 negative control + done. 115 let wrong: i64=f32_mul(dx[0], f32_of(2)) 116 let fd0: i64=f32_div(f32_sub(rms_L(x,gamma,w,N,eps), rms_L(x,gamma,w,N,eps)), twoh) // 0 (placeholder), real check below 117 var negok: i64=0 118 let xp0: *i64=sys_mmap(128) as *i64; let xm0: *i64=sys_mmap(128) as *i64; var c4: i64=0 119 while c4<N { xp0[c4]=x[c4]; xm0[c4]=x[c4]; c4=c4+1 } 120 xp0[0]=f32_add(x[0],h); xm0[0]=f32_sub(x[0],h) 121 let realfd0: i64=f32_div(f32_sub(rms_L(xp0,gamma,w,N,eps), rms_L(xm0,gamma,w,N,eps)), twoh) 122 if f32_le(f32_abs(f32_sub(wrong, realfd0)), tol)==0 { negok=1 } 123 total=total+1; if gok==1 { if ropeok==1 { if negok==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) } 124 gw("T5 NEGATIVE CONTROL + DONE: a doubled RMSNorm grad " as *u8); gm(wrong); gw("m REJECTED vs fd " as *u8); gm(realfd0); gw("m -> Llama-arch nonlinearities (RMSNorm+RoPE) backprop in f32\n" as *u8) 125 126 gw("\n RUNG 3 DONE: the Llama-3.2 lightweight nonlinearities backprop correctly in f32 -- RMSNorm (gradchecked, scale-invariant)\n" as *u8) 127 gw(" + RoPE (rotation backward = inverse rotation, gradchecked) + f32_sin/cos. HONEST: RoPE sin/cos here are small-angle Taylor;\n" as *u8) 128 gw(" production needs range-reduction (mod 2pi) for large position*theta. With softmax (R2a), the attention+norm stack is covered.\n" as *u8) 129 gw(" NEXT: GQA wiring (grouped-query attention) + SwiGLU (SiLU MLP), then R3 Adam, R4 data (crawl-on-burst), R5 train loop.\n" as *u8) 130 gw("F32-RMSNORM-ROPE verdict=" as *u8) 131 if pass==total { gw("GREEN passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw(" -- RMSNorm + RoPE backward correct (gradchecked), Llama-arch nonlinearities\n" as *u8); sys_exit(0); return 0 } 132 gw("RED passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw("\n" as *u8); sys_exit(1); return 1 133}