code wiki / _hdl_build / nx_f32_layernorm_gate.nx
nx_f32_layernorm_gate.nx source
↩ module page · 134 lines · 9683 B
1import "nx_gate_gn.nx"
2import "nx_gate_base.nx"
3// nx_f32_layernorm_gate.nx -- RUNG 2b of the from-scratch 0.5B training loop: LAYERNORM forward & backward, gradient-
4// checked. Layernorm is the transformer's other nonlinearity (softmax was R2a). Its backward is the hard one: mean and
5// variance depend on ALL inputs, so dL/dx_i couples through the whole vector -- the finite-difference check is the real
6// proof the coupling is right. Uses f32_sqrt (built+tested in R2a) + the R1 f32 autograd discipline. y_i = gamma_i *
7// xhat_i + beta_i where xhat = (x-mean)/sqrt(var+eps); learnable gamma,beta get gradients too.
8// T0 FORWARD: xhat is mean-0 unit-variance; layernorm([1,2,3]) -> xhat=[-1.225,0,1.225].
9// T1 dL/dx GRADCHECK: the coupled normalization backward == central finite-difference (the hard proof).
10// T2 SHIFT-INVARIANCE: sum_i dL/dx_i = 0 (shifting all x by c doesn't change layernorm -> gradients sum to zero).
11// T3 dL/dgamma, dL/dbeta GRADCHECK: the learnable scale/shift gradients == finite-difference.
12// T4 NEGATIVE CONTROL: a doubled (wrong) dL/dx is REJECTED.
13// T5 = layernorm backward correct -> with softmax (R2a), the full transformer block backprops in f32. R2 DONE.
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)))) }
21
22func f32_abs(x: i64) -> i64 { return x & 0x7FFFFFFF }
23func 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 }
24// f32_sqrt via Newton (same as R2a; lives in a shared f32-math lib once the train loop is the 3rd user).
25func 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 }
26
27// LAYERNORM forward: stores xhat + sigma for the backward. eps stabilizes sqrt(var).
28func layernorm_fwd(x: *i64, gamma: *i64, beta: *i64, n: i64, eps: i64, y: *i64, xhat: *i64, sig: *i64) -> i64 {
29 var sum: i64=f32_of(0); var i: i64=0; while i<n { sum=f32_add(sum, x[i]); i=i+1 }
30 let mu: i64=f32_div(sum, f32_of(n))
31 var vs: i64=f32_of(0); i=0; while i<n { let d: i64=f32_sub(x[i],mu); vs=f32_add(vs, f32_mul(d,d)); i=i+1 }
32 let sigma: i64=f32_sqrt(f32_add(f32_div(vs, f32_of(n)), eps)); sig[0]=sigma
33 i=0; while i<n { xhat[i]=f32_div(f32_sub(x[i],mu), sigma); y[i]=f32_add(f32_mul(gamma[i],xhat[i]), beta[i]); i=i+1 }
34 return 0
35}
36// LAYERNORM backward: dL/dx_i = (1/sigma)*(dxhat_i - mean(dxhat) - xhat_i*mean(dxhat*xhat)), dxhat=g*gamma.
37func layernorm_bwd(g: *i64, gamma: *i64, xhat: *i64, sigma: i64, n: i64, dx: *i64) -> i64 {
38 let dxhat: *i64=sys_mmap(128) as *i64
39 var s1: i64=f32_of(0); var s2: i64=f32_of(0); var i: i64=0
40 while i<n { dxhat[i]=f32_mul(g[i],gamma[i]); s1=f32_add(s1,dxhat[i]); s2=f32_add(s2,f32_mul(dxhat[i],xhat[i])); i=i+1 }
41 let mdx: i64=f32_div(s1, f32_of(n)); let mdxx: i64=f32_div(s2, f32_of(n))
42 let invs: i64=f32_div(f32_of(1), sigma)
43 i=0; while i<n { dx[i]=f32_mul(invs, f32_sub(f32_sub(dxhat[i], mdx), f32_mul(xhat[i], mdxx))); i=i+1 }
44 return 0
45}
46// L = sum_i w_i * y_i (forward, for the finite-difference probe).
47func compute_L(x: *i64, gamma: *i64, beta: *i64, w: *i64, n: i64, eps: i64) -> i64 {
48 let y: *i64=sys_mmap(128) as *i64; let xh: *i64=sys_mmap(128) as *i64; let sg: *i64=sys_mmap(16) as *i64
49 layernorm_fwd(x, gamma, beta, n, eps, y, xh, sg)
50 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 }
51 return L
52}
53
54func main() -> i64 {
55 gw("=== nx_f32_layernorm_gate: RUNG 2b -- layernorm forward/backward (gradchecked), the coupled normalization backprop ===\n" as *u8)
56 var pass: i64=0; var total: i64=0
57 let eps: i64=f32_div(f32_of(1), f32_of(100000)) // 1e-5
58 let h: i64=f32_div(f32_of(1), f32_of(100)) // 0.01 finite-diff step
59 let tol: i64=f32_div(f32_of(3), f32_of(100)) // 0.03
60 let twoh: i64=f32_mul(f32_of(2), h)
61
62 let x: *i64=sys_mmap(128) as *i64; x[0]=f32_of(1); x[1]=f32_of(2); x[2]=f32_of(3)
63 let gamma: *i64=sys_mmap(128) as *i64; gamma[0]=f32_of(1); gamma[1]=f32_of(1); gamma[2]=f32_of(1)
64 let beta: *i64=sys_mmap(128) as *i64; beta[0]=f32_of(0); beta[1]=f32_of(0); beta[2]=f32_of(0)
65 let w: *i64=sys_mmap(128) as *i64; w[0]=f32_of(1); w[1]=f32_of(0); w[2]=f32_of(0) // L = y_0
66 let N: i64=3
67 let y: *i64=sys_mmap(128) as *i64; let xhat: *i64=sys_mmap(128) as *i64; let sg: *i64=sys_mmap(16) as *i64
68 layernorm_fwd(x, gamma, beta, N, eps, y, xhat, sg)
69
70 // T0 FORWARD.
71 total=total+1; let xh0: i64=f32_int(f32_mul(xhat[0],f32_of(1000))); let xh2: i64=f32_int(f32_mul(xhat[2],f32_of(1000)))
72 if xh0<=-1220 { if xh0>=-1230 { if xh2>=1220 { if xh2<=1230 { 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) }
73 gw("T0 FORWARD: layernorm([1,2,3]) xhat=[" as *u8); gm(xhat[0]); gw("," as *u8); gm(xhat[1]); gw("," as *u8); gm(xhat[2]); gw("]m (expect -1225,0,1225)\n" as *u8)
74
75 // T1 dL/dx GRADCHECK.
76 let dx: *i64=sys_mmap(128) as *i64; layernorm_bwd(w, gamma, xhat, sg[0], N, dx)
77 var gok: i64=1; var j: i64=0
78 while j<N {
79 let xp: *i64=sys_mmap(128) as *i64; let xm: *i64=sys_mmap(128) as *i64; var c: i64=0
80 while c<N { xp[c]=x[c]; xm[c]=x[c]; c=c+1 }
81 xp[j]=f32_add(x[j],h); xm[j]=f32_sub(x[j],h)
82 let fd: i64=f32_div(f32_sub(compute_L(xp,gamma,beta,w,N,eps), compute_L(xm,gamma,beta,w,N,eps)), twoh)
83 if f32_le(f32_abs(f32_sub(dx[j], fd)), tol)==0 { gok=0 }
84 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)
85 j=j+1
86 }
87 total=total+1; if gok==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
88 gw("T1 dL/dx GRADCHECK: coupled normalization backward == finite-diff (the hard proof: mean/var couple all inputs)\n" as *u8)
89
90 // T2 SHIFT-INVARIANCE: sum dL/dx ~= 0.
91 let dxsum: i64=f32_int(f32_mul(f32_add(f32_add(dx[0],dx[1]),dx[2]), f32_of(1000)))
92 total=total+1; if dxsum>=-3 { if dxsum<=3 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) }
93 gw("T2 SHIFT-INVARIANCE: sum(dL/dx)=" as *u8); gn(dxsum); gw("m ~= 0 (shifting x by a constant leaves layernorm unchanged)\n" as *u8)
94
95 // T3 dL/dgamma + dL/dbeta GRADCHECK. analytic: dL/dgamma_i = w_i*xhat_i ; dL/dbeta_i = w_i.
96 var pgok: i64=1
97 j=0; while j<N {
98 let gp: *i64=sys_mmap(128) as *i64; let gmn: *i64=sys_mmap(128) as *i64; var c: i64=0
99 while c<N { gp[c]=gamma[c]; gmn[c]=gamma[c]; c=c+1 }
100 gp[j]=f32_add(gamma[j],h); gmn[j]=f32_sub(gamma[j],h)
101 let fdg: i64=f32_div(f32_sub(compute_L(x,gp,beta,w,N,eps), compute_L(x,gmn,beta,w,N,eps)), twoh)
102 let ana_g: i64=f32_mul(w[j], xhat[j])
103 if f32_le(f32_abs(f32_sub(ana_g, fdg)), tol)==0 { pgok=0 }
104 let bp: *i64=sys_mmap(128) as *i64; let bmn: *i64=sys_mmap(128) as *i64; c=0
105 while c<N { bp[c]=beta[c]; bmn[c]=beta[c]; c=c+1 }
106 bp[j]=f32_add(beta[j],h); bmn[j]=f32_sub(beta[j],h)
107 let fdb: i64=f32_div(f32_sub(compute_L(x,gamma,bp,w,N,eps), compute_L(x,gamma,bmn,w,N,eps)), twoh)
108 if f32_le(f32_abs(f32_sub(w[j], fdb)), tol)==0 { pgok=0 }
109 j=j+1
110 }
111 total=total+1; if pgok==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
112 gw("T3 dL/dgamma + dL/dbeta GRADCHECK: learnable scale/shift gradients == finite-diff\n" as *u8)
113
114 // T4 NEGATIVE CONTROL.
115 let xp0: *i64=sys_mmap(128) as *i64; let xm0: *i64=sys_mmap(128) as *i64; var c3: i64=0
116 while c3<N { xp0[c3]=x[c3]; xm0[c3]=x[c3]; c3=c3+1 }
117 xp0[0]=f32_add(x[0],h); xm0[0]=f32_sub(x[0],h)
118 let fd0: i64=f32_div(f32_sub(compute_L(xp0,gamma,beta,w,N,eps), compute_L(xm0,gamma,beta,w,N,eps)), twoh)
119 let wrong: i64=f32_mul(dx[0], f32_of(2))
120 total=total+1; if f32_le(f32_abs(f32_sub(wrong, fd0)), tol)==0 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) }
121 gw("T4 NEGATIVE CONTROL: a doubled dL/dx0 " as *u8); gm(wrong); gw("m is REJECTED vs finite-diff " as *u8); gm(fd0); gw("m\n" as *u8)
122
123 // T5.
124 total=total+1; if gok==1 { if pgok==1 { pass=pass+1; gw(" [PASS] " as *u8) } else { gw(" [FAIL] " as *u8) } } else { gw(" [FAIL] " as *u8) }
125 gw("T5 R2 DONE: softmax (R2a) + layernorm (R2b) backward both gradchecked -> the full transformer block backprops in f32\n" as *u8)
126
127 gw("\n RUNG 2 COMPLETE: both transformer nonlinearities backprop correctly in f32 -- softmax (attention) and layernorm, each\n" as *u8)
128 gw(" finite-difference verified with a working negative control. With R1 (autograd) + R2 (softmax+layernorm), a transformer\n" as *u8)
129 gw(" block's gradients are covered. NEXT: R3 Adam optimizer (m,v moments + bias-correct), R4 tokenized data pipeline (fed by\n" as *u8)
130 gw(" crawl-on-burst), R5 the train loop -> train the from-scratch 0.5B (~$50 cloud / ~1 week on the RTX 5080 / ~1 day on 1 H100).\n" as *u8)
131 gw("F32-LAYERNORM verdict=" as *u8)
132 if pass==total { gw("GREEN passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw(" -- layernorm backward correct (gradchecked); R2 complete\n" as *u8); sys_exit(0); return 0 }
133 gw("RED passes=" as *u8); gn(pass); gw("/" as *u8); gn(total); gw("\n" as *u8); sys_exit(1); return 1
134}