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}