code wiki / (root) / nx_groupnorm.nx

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}