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}