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}