code wiki / (root) / nx_f32_layernorm.nx

nx_f32_layernorm.nx source

↩ module page · 34 lines · 1484 B

1// nx_f32_layernorm.nx -- clean importable flat-array software-f32 LayerNorm (per-row normalize over `dim`, then 2// affine): out = (x-mean)/sqrt(var+eps) * gamma + beta. The ecosystem has nx_layernorm_forward (NxTensor) and a 3// layernorm_fwd embedded in a gate, but no plain importable flat lib -- this is it, the ViT/BERT encoder building 4// block. Composes nx_f32_add/sub/mul/div/sqrt. license_tier: ORIGINAL 5import "nx_syscalls.nx" 6import "nx_f32.nx" 7import "nx_f32_cvt.nx" 8import "nx_f32_div.nx" 9 10// x,out are [n_rows, dim] row-major; gamma,beta are [dim]; eps is an f32 scalar. 11func nx_f32_layernorm(x: *i64, gamma: *i64, beta: *i64, n_rows: i64, dim: i64, eps: i64, out: *i64) -> i64 { 12 let dimf: i64 = nx_i32_to_f32(dim) 13 var r: i64 = 0 14 while r < n_rows { 15 let base: i64 = r * dim 16 var s: i64 = 0 17 var i: i64 = 0 18 while i < dim { s = nx_f32_add(s, x[base + i]); i = i + 1 } 19 let mean: i64 = nx_f32_div(s, dimf) 20 var v: i64 = 0 21 i = 0 22 while i < dim { let d: i64 = nx_f32_sub(x[base + i], mean); v = nx_f32_add(v, nx_f32_mul(d, d)); i = i + 1 } 23 let varr: i64 = nx_f32_div(v, dimf) 24 let denom: i64 = nx_f32_sqrt(nx_f32_add(varr, eps)) 25 i = 0 26 while i < dim { 27 let xhat: i64 = nx_f32_div(nx_f32_sub(x[base + i], mean), denom) 28 out[base + i] = nx_f32_add(nx_f32_mul(xhat, gamma[i]), beta[i]) 29 i = i + 1 30 } 31 r = r + 1 32 } 33 return 0 34}