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}