nx_groupnorm.nx source
↩ module page · 324 lines · 12424 B
1// nx_groupnorm.nx -- Group Normalisation (Wu & He 2018).
2//
3// Closes the missing-dep gap for the UNet block. Modern image-gen
4// architectures use GroupNorm, not RMSNorm or LayerNorm:
5//
6// Stable Diffusion 1.x / 2.x / 3 UNet: GroupNorm(num_groups=32)
7// Flux UNet: GroupNorm
8// Z-Image: GroupNorm
9// StyleGAN family: GroupNorm + AdaIN (queued)
10//
11// Substrate had RMSNorm (decoder-only LLMs) + LayerNorm (encoder-
12// decoders / GPT-2 era). GroupNorm was the third missing norm
13// variant and the load-bearing one for diffusion.
14//
15// ===== Math (Wu & He 2018 _Group Normalization_) =================
16//
17// Input: x [N, C, H, W] Q10
18// Groups: G (C must be divisible by G; typically G=32)
19//
20// For each (n, g):
21// mean = mean of x over (C/G channels in group g, H pixels, W pixels)
22// (i.e., (C/G * H * W) elements per group per sample)
23// var = variance over the same elements
24// For each (c in group g, h, w):
25// y[n, c, h, w] = (x[n, c, h, w] - mean) / sqrt(var + eps) *
26// gamma[c] + beta[c]
27//
28// gamma is the per-channel learned scale (length C, Q10).
29// beta is the per-channel learned bias (length C, Q10).
30//
31// Group-count edge cases:
32// G = 1 -> LayerNorm-like (normalize over all C*H*W)
33// G = C -> InstanceNorm (normalize per-channel separately)
34//
35// Standard GroupNorm uses G=32 (Wu & He's recommendation, picked
36// to be roughly invariant across batch size).
37//
38// ===== Q-format ===================================================
39//
40// x, gamma, beta all in Q10 (substrate convention). Accumulators
41// in raw integer (Q20 for sum_sq), divided to Q10 before sqrt.
42// eps_q10 = 1 (tightest Q10 stabiliser).
43//
44// Per the bits-up cardinal + bounded-loop discipline.
45//
46// genealogy_id: wu_he_2018_group_norm + ba_kiros_hinton_2016_layer_norm +
47// ulyanov_2017_instance_norm + ioffe_szegedy_2015_batch_norm
48// lineage_id: substrate_groupnorm_v1
49
50// nx_safety_envelope:
51// intended_use: AUTO_APPLIED -- primitive-specific tuning queued
52// sil_target: SIL1
53// evidence: [bulk_applied_2026-05-16, see-file-comment-for-detail]
54// verdict: NOT_YET_EVALUATED
55
56import "nx_syscalls.nx"
57import "nx_tier.nx"
58import "nx_loop.nx"
59import "nx_tensor.nx"
60import "nx_isqrt.nx"
61
62// ===== Constants ==================================================
63
64const NX_GN_Q10: nx_int = 1024
65const NX_GN_EPS_Q10: nx_int = 1
66
67// ===== Sealed-enum: GroupNormVerdict ==============================
68
69const NX_GN_OK: nx_int = 0
70const NX_GN_ERR_BAD_DTYPE: nx_int = 1
71const NX_GN_ERR_BAD_NDIM: nx_int = 2
72const NX_GN_ERR_SHAPE_MISMATCH: nx_int = 3
73const NX_GN_ERR_NOT_CONTIGUOUS: nx_int = 4
74const NX_GN_ERR_BAD_GROUPS: nx_int = 5
75const NX_GN_N_VERDICTS: nx_int = 6
76
77func nx_gn_verdict_is_valid(v: nx_int) -> nx_int {
78 if v < 0 { return 0 }
79 if v >= NX_GN_N_VERDICTS { return 0 }
80 return 1
81}
82
83// ===== Forward pass ==============================================
84//
85// x: *NxTensor [N, C, H, W] Q10 input
86// G: number of groups (must divide C cleanly)
87// gamma: *i64 [C] Q10 per-channel scale
88// beta: *i64 [C] Q10 per-channel bias (nullable; 0 = no bias)
89// out: *NxTensor [N, C, H, W] Q10 output (in-place supported)
90
91func nx_groupnorm_forward(x: *NxTensor, n_groups: nx_int,
92 gamma: *i64, beta: *i64,
93 out: *NxTensor) -> nx_int {
94 if x.dtype != NX_DT_I64 { return NX_GN_ERR_BAD_DTYPE }
95 if out.dtype != NX_DT_I64 { return NX_GN_ERR_BAD_DTYPE }
96 if x.ndim != 4 { return NX_GN_ERR_BAD_NDIM }
97 if out.ndim != 4 { return NX_GN_ERR_BAD_NDIM }
98 if x.shape[0] != out.shape[0] { return NX_GN_ERR_SHAPE_MISMATCH }
99 if x.shape[1] != out.shape[1] { return NX_GN_ERR_SHAPE_MISMATCH }
100 if x.shape[2] != out.shape[2] { return NX_GN_ERR_SHAPE_MISMATCH }
101 if x.shape[3] != out.shape[3] { return NX_GN_ERR_SHAPE_MISMATCH }
102 if nx_t_is_contiguous(x) == 0 { return NX_GN_ERR_NOT_CONTIGUOUS }
103 if nx_t_is_contiguous(out) == 0 { return NX_GN_ERR_NOT_CONTIGUOUS }
104
105 let N: nx_int = x.shape[0]
106 let C: nx_int = x.shape[1]
107 let H: nx_int = x.shape[2]
108 let W: nx_int = x.shape[3]
109
110 if n_groups <= 0 { return NX_GN_ERR_BAD_GROUPS }
111 if C - (C / n_groups) * n_groups != 0 { return NX_GN_ERR_BAD_GROUPS }
112 let C_per_group: nx_int = C / n_groups
113 let group_size: nx_int = C_per_group * H * W
114 let chan_stride: nx_int = H * W
115 let batch_stride: nx_int = C * H * W
116
117 let px: *i64 = x.storage as *i64
118 let po: *i64 = out.storage as *i64
119
120 var n: nx_int = 0
121 var n_iter: nx_int = 0
122 var n_verdict: nx_int = NX_LOOP_RUNNING
123 let N_BUDGET: nx_int = N
124 while n_verdict == NX_LOOP_RUNNING && n_iter < N_BUDGET {
125 var g: nx_int = 0
126 var g_iter: nx_int = 0
127 var g_verdict: nx_int = NX_LOOP_RUNNING
128 let G_BUDGET: nx_int = n_groups
129 while g_verdict == NX_LOOP_RUNNING && g_iter < G_BUDGET {
130 // Channel range for this group.
131 let c_start: nx_int = g * C_per_group
132 let c_end: nx_int = c_start + C_per_group
133
134 // Pass 1: mean across group_size elements.
135 var sum: i64 = 0
136 var c: nx_int = c_start
137 var c_iter: nx_int = 0
138 var c_verdict: nx_int = NX_LOOP_RUNNING
139 while c_verdict == NX_LOOP_RUNNING && c_iter < C_per_group {
140 let chan_base: nx_int = n * batch_stride + c * chan_stride
141 var p: nx_int = 0
142 var p_iter: nx_int = 0
143 var p_verdict: nx_int = NX_LOOP_RUNNING
144 let P_BUDGET: nx_int = chan_stride
145 while p_verdict == NX_LOOP_RUNNING && p_iter < P_BUDGET {
146 sum = sum + px[chan_base + p]
147 p = p + 1
148 p_iter = p_iter + 1
149 }
150 c = c + 1
151 c_iter = c_iter + 1
152 }
153 let mean: i64 = sum / group_size
154
155 // Pass 2: variance (centred sum of squares).
156 var sum_sq: i64 = 0
157 var c2: nx_int = c_start
158 var c2_iter: nx_int = 0
159 var c2_verdict: nx_int = NX_LOOP_RUNNING
160 while c2_verdict == NX_LOOP_RUNNING && c2_iter < C_per_group {
161 let chan_base: nx_int = n * batch_stride + c2 * chan_stride
162 var p: nx_int = 0
163 var p_iter: nx_int = 0
164 var p_verdict: nx_int = NX_LOOP_RUNNING
165 while p_verdict == NX_LOOP_RUNNING && p_iter < chan_stride {
166 let dx: i64 = px[chan_base + p] - mean
167 sum_sq = sum_sq + dx * dx
168 p = p + 1
169 p_iter = p_iter + 1
170 }
171 c2 = c2 + 1
172 c2_iter = c2_iter + 1
173 }
174 let var_q10: i64 = sum_sq / (group_size * NX_GN_Q10)
175 let std_q10: i64 = nx_isqrt_q10(var_q10 + NX_GN_EPS_Q10)
176 if std_q10 <= 0 { g_verdict = NX_LOOP_DONE_EXIT }
177
178 // Pass 3: normalise + per-channel affine.
179 if g_verdict == NX_LOOP_RUNNING {
180 var c3: nx_int = c_start
181 var c3_iter: nx_int = 0
182 var c3_verdict: nx_int = NX_LOOP_RUNNING
183 while c3_verdict == NX_LOOP_RUNNING && c3_iter < C_per_group {
184 let chan_base: nx_int = n * batch_stride + c3 * chan_stride
185 let g_val: i64 = gamma[c3]
186 var b_val: i64 = 0
187 if (beta as i64) != 0 { b_val = beta[c3] }
188 var p: nx_int = 0
189 var p_iter: nx_int = 0
190 var p_verdict: nx_int = NX_LOOP_RUNNING
191 while p_verdict == NX_LOOP_RUNNING && p_iter < chan_stride {
192 let centred: i64 = px[chan_base + p] - mean
193 let norm: i64 = (centred * NX_GN_Q10) / std_q10
194 po[chan_base + p] = (norm * g_val) / NX_GN_Q10 + b_val
195 p = p + 1
196 p_iter = p_iter + 1
197 }
198 c3 = c3 + 1
199 c3_iter = c3_iter + 1
200 }
201 }
202 g = g + 1
203 g_iter = g_iter + 1
204 }
205 n = n + 1
206 n_iter = n_iter + 1
207 }
208 return NX_GN_OK
209}
210
211// ===== Factory: unit-affine gamma + zero beta ====================
212
213func nx_groupnorm_gamma_unit(c: nx_int) -> *i64 {
214 let g: *i64 = sys_mmap(c * 8) as *i64
215 var i: nx_int = 0
216 var iter: nx_int = 0
217 var verdict: nx_int = NX_LOOP_RUNNING
218 let BUDGET: nx_int = c
219 while verdict == NX_LOOP_RUNNING && iter < BUDGET {
220 g[i] = NX_GN_Q10
221 i = i + 1
222 iter = iter + 1
223 }
224 return g
225}
226
227func nx_groupnorm_beta_zero(c: nx_int) -> *i64 {
228 let b: *i64 = sys_mmap(c * 8) as *i64
229 var i: nx_int = 0
230 var iter: nx_int = 0
231 var verdict: nx_int = NX_LOOP_RUNNING
232 let BUDGET: nx_int = c
233 while verdict == NX_LOOP_RUNNING && iter < BUDGET {
234 b[i] = 0
235 i = i + 1
236 iter = iter + 1
237 }
238 return b
239}
240
241// ===== Self-test ==================================================
242//
243// Closed-form invariants:
244//
245// (a) Constant input: var=0 + eps stabilisation -> output ~0 per group
246// (b) Shift invariance: GroupNorm(x + k) = GroupNorm(x)
247// (c) Bad group count (doesn't divide C) -> ERR_BAD_GROUPS
248// (d) Verdict-range gate
249
250func main() -> i64 {
251 let N: nx_int = 1
252 let C: nx_int = 4 // 2 groups of 2 channels each
253 let H: nx_int = 2
254 let W: nx_int = 2
255
256 let sh: *nx_int = sys_mmap(4 * 8) as *nx_int
257 sh[0]=N; sh[1]=C; sh[2]=H; sh[3]=W
258
259 let err: *nx_int = sys_mmap(8) as *nx_int
260 err[0] = 0
261 let xt: *NxTensor = nx_t_alloc(NX_DT_I64, sh, 4, err)
262 let yt: *NxTensor = nx_t_alloc(NX_DT_I64, sh, 4, err)
263 if err[0] != 0 { return 5 }
264
265 let gamma: *i64 = nx_groupnorm_gamma_unit(C)
266 let beta: *i64 = nx_groupnorm_beta_zero(C)
267
268 // --- (a) Constant tensor ---
269 let px: *i64 = xt.storage as *i64
270 let py: *i64 = yt.storage as *i64
271 var i: nx_int = 0
272 while i < N * C * H * W { px[i] = 5 * NX_GN_Q10; i = i + 1 }
273
274 let v_a: nx_int = nx_groupnorm_forward(xt, 2, gamma, beta, yt)
275 if v_a != NX_GN_OK { return 10 + v_a }
276 // After normalising a constant input + eps stabilisation, output
277 // should be small (close to 0).
278 var ai: nx_int = 0
279 while ai < N * C * H * W {
280 if py[ai] > 200 { return 20 }
281 if py[ai] < -200 { return 21 }
282 ai = ai + 1
283 }
284
285 // --- (b) Shift invariance ---
286 // Fill with two different patterns per group, run both, verify
287 // the per-group outputs are equal regardless of the constant
288 // offset applied to one of them.
289 // Channel 0: 100, 200, 300, 400.
290 // Channel 1: 500, 600, 700, 800.
291 // Channel 2: 100+k, 200+k, 300+k, 400+k. (group 1, shifted)
292 // Channel 3: 500+k, 600+k, 700+k, 800+k.
293 let K: i64 = 1000 * NX_GN_Q10
294 px[0]=100*NX_GN_Q10; px[1]=200*NX_GN_Q10; px[2]=300*NX_GN_Q10; px[3]=400*NX_GN_Q10
295 px[4]=500*NX_GN_Q10; px[5]=600*NX_GN_Q10; px[6]=700*NX_GN_Q10; px[7]=800*NX_GN_Q10
296 px[8]=100*NX_GN_Q10+K; px[9]=200*NX_GN_Q10+K; px[10]=300*NX_GN_Q10+K; px[11]=400*NX_GN_Q10+K
297 px[12]=500*NX_GN_Q10+K; px[13]=600*NX_GN_Q10+K; px[14]=700*NX_GN_Q10+K; px[15]=800*NX_GN_Q10+K
298
299 nx_groupnorm_forward(xt, 2, gamma, beta, yt)
300 // py[0..8] -> group 0 (channels 0, 1)
301 // py[8..16] -> group 1 (channels 2, 3, shifted-input)
302 // Group 1's input is GROUP 0's input + K applied uniformly; thus
303 // mean shifts by K but normalised output equals group 0's.
304 var bi: nx_int = 0
305 while bi < 8 {
306 let drift: i64 = py[8 + bi] - py[bi]
307 if drift > 8 { return 30 }
308 if drift < -8 { return 31 }
309 bi = bi + 1
310 }
311
312 // --- (c) Bad groups count ---
313 let v_bad: nx_int = nx_groupnorm_forward(xt, 3, gamma, beta, yt)
314 if v_bad != NX_GN_ERR_BAD_GROUPS { return 40 } // 4 / 3 doesn't divide
315
316 // --- (d) Verdict gate ---
317 var vi: nx_int = 0
318 while vi < NX_GN_N_VERDICTS {
319 if nx_gn_verdict_is_valid(vi) != 1 { return 50 + vi }
320 vi = vi + 1
321 }
322
323 return 0
324}