code wiki / (root) / nx_f32_layernorm.nx

nx_f32_layernorm.nx source

↩ module page · 63 lines · 2780 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// ---- HARDWARE TWIN (2026-09-14, search R0 cross-encoder): the same arithmetic on the __f32 intrinsics; subtraction 11// is the sign-flipped add (a + (-b), exactly what nx_f32_sub computes), the one square root per row stays software. 12// DIFFERENTIALLY GATED bit for bit against nx_f32_layernorm_sw below by nx_f32_bricks_hw_gate. 13func nx_f32_layernorm(x: *i64, gamma: *i64, beta: *i64, n_rows: i64, dim: i64, eps: i64, out: *i64) -> i64 { 14 let dimf: i64 = nx_i32_to_f32(dim) 15 var r: i64 = 0 16 while r < n_rows { 17 let base: i64 = r * dim 18 var s: i64 = 0 19 var i: i64 = 0 20 while i < dim { s = __f32_add(s, x[base + i]); i = i + 1 } 21 let mean: i64 = __f32_div(s, dimf) 22 let nmean: i64 = mean ^ 0x80000000 23 var v: i64 = 0 24 i = 0 25 while i < dim { let d: i64 = __f32_add(x[base + i], nmean); v = __f32_add(v, __f32_mul(d, d)); i = i + 1 } 26 let varr: i64 = __f32_div(v, dimf) 27 let denom: i64 = nx_f32_sqrt(__f32_add(varr, eps)) 28 i = 0 29 while i < dim { 30 let xhat: i64 = __f32_div(__f32_add(x[base + i], nmean), denom) 31 out[base + i] = __f32_add(__f32_mul(xhat, gamma[i]), beta[i]) 32 i = i + 1 33 } 34 r = r + 1 35 } 36 return 0 37} 38 39// x,out are [n_rows, dim] row-major; gamma,beta are [dim]; eps is an f32 scalar. THE ORACLE, kept verbatim. 40func nx_f32_layernorm_sw(x: *i64, gamma: *i64, beta: *i64, n_rows: i64, dim: i64, eps: i64, out: *i64) -> i64 { 41 let dimf: i64 = nx_i32_to_f32(dim) 42 var r: i64 = 0 43 while r < n_rows { 44 let base: i64 = r * dim 45 var s: i64 = 0 46 var i: i64 = 0 47 while i < dim { s = nx_f32_add(s, x[base + i]); i = i + 1 } 48 let mean: i64 = nx_f32_div(s, dimf) 49 var v: i64 = 0 50 i = 0 51 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 } 52 let varr: i64 = nx_f32_div(v, dimf) 53 let denom: i64 = nx_f32_sqrt(nx_f32_add(varr, eps)) 54 i = 0 55 while i < dim { 56 let xhat: i64 = nx_f32_div(nx_f32_sub(x[base + i], mean), denom) 57 out[base + i] = nx_f32_add(nx_f32_mul(xhat, gamma[i]), beta[i]) 58 i = i + 1 59 } 60 r = r + 1 61 } 62 return 0 63}