nx_gen_prenorm_mod_verify.nx source
↩ module page · 161 lines · 6207 B
1// nx_gen_prenorm_mod_verify.nx -- SOVEREIGN DiT pre-norm + adaLN modulate, verified vs the oracle.
2//
3// y[t][i] = ( x[t][i] / sqrt(mean_i(x[t]^2) + eps) * w[i] ) * (1 + scale[i])
4//
5// This is the entry half of every modulated DiT block: RMSNorm, then scale-modulation broadcast
6// from one chunk of the adaLN vector. Z-Image, Flux and SD3/MMDiT all use this shape; only the
7// chunk COUNT and INDEX differ, so both are arguments.
8//
9// Usage:
10// nx_gen_prenorm_mod_verify <model> <x> <norm_w> <adaln> <n_chunks> <scale_idx> <y> [rows] [eps_recip]
11//
12// It is a TWO-op composite on purpose. RMSNorm is already verified standalone by
13// nx_gen_rmsnorm_verify, so if this organ is RED while that one is GREEN the fault is isolated to
14// the modulate half or to the chunk index -- which is the only reason a composite is acceptable
15// here rather than a fixture for the intermediate.
16//
17// Hardware __f32_* intrinsics in the hot loop, never the nx_f32_* software twins.
18// license_tier: ORIGINAL
19
20import "nx_syscalls.nx"
21import "nx_le.nx"
22import "nx_f32.nx"
23import "nx_f32_div.nx"
24import "nx_f32_cvt.nx"
25import "nx_strconv.nx"
26import "nx_genfix.nx"
27import "nx_genver.nx"
28const K_MAGIC_1000000: i64 = 1000000
29
30func zp_strlen(s: *u8) -> i64 {
31 var n: i64 = 0
32 while s[n] != (0 as u8) { n = n + 1 }
33 return n
34}
35
36func main(argc: i64, argv: *i64) -> i64 {
37 if argc < 8 {
38 nx_genver_emit("usage_model_x_normw_adaln_nchunks_scaleidx_y" as *u8, argc)
39 return 2
40 }
41 let model: *u8 = argv[1] as *u8
42 let xn: *u8 = argv[2] as *u8
43 let wn: *u8 = argv[3] as *u8
44 let an: *u8 = argv[4] as *u8
45 let errp: *i64 = sys_mmap(32) as *i64
46 errp[0] = 0
47 let n_chunks: i64 = nx_strconv_parse_i64(argv[5] as *u8, errp)
48 if errp[0] != 0 { nx_genver_emit("bad_n_chunks" as *u8, 1); return 3 }
49 errp[0] = 0
50 let scale_idx: i64 = nx_strconv_parse_i64(argv[6] as *u8, errp)
51 if errp[0] != 0 { nx_genver_emit("bad_scale_idx" as *u8, 1); return 4 }
52 let yn: *u8 = argv[7] as *u8
53
54 var rows: i64 = 4
55 if argc >= 9 {
56 errp[0] = 0
57 rows = nx_strconv_parse_i64(argv[8] as *u8, errp)
58 if errp[0] != 0 { nx_genver_emit("bad_rows" as *u8, 1); return 5 }
59 }
60 var eps_recip: i64 = K_MAGIC_1000000
61 if argc >= 10 {
62 errp[0] = 0
63 eps_recip = nx_strconv_parse_i64(argv[9] as *u8, errp)
64 if errp[0] != 0 { nx_genver_emit("bad_eps_recip" as *u8, 1); return 6 }
65 }
66 if n_chunks <= 0 { nx_genver_emit("bad_n_chunks" as *u8, n_chunks); return 7 }
67 if scale_idx < 0 || scale_idx >= n_chunks {
68 nx_genver_emit("scale_idx_out_of_range" as *u8, scale_idx)
69 return 8
70 }
71
72 let lx: i64 = zp_strlen(xn)
73 let lw: i64 = zp_strlen(wn)
74 let la: i64 = zp_strlen(an)
75 let ly: i64 = zp_strlen(yn)
76
77 let ne_x: *i64 = sys_mmap(64) as *i64
78 let ne_w: *i64 = sys_mmap(64) as *i64
79 let ne_a: *i64 = sys_mmap(64) as *i64
80 let ne_y: *i64 = sys_mmap(64) as *i64
81 let c_x: i64 = nx_genfix_dims(model, xn, lx, ne_x)
82 let c_w: i64 = nx_genfix_dims(model, wn, lw, ne_w)
83 let c_a: i64 = nx_genfix_dims(model, an, la, ne_a)
84 let c_y: i64 = nx_genfix_dims(model, yn, ly, ne_y)
85 if c_x < 0 { nx_genver_emit("missing_x" as *u8, 1); return 30 }
86 if c_w < 0 { nx_genver_emit("missing_norm_w" as *u8, 1); return 31 }
87 if c_a < 0 { nx_genver_emit("missing_adaln" as *u8, 1); return 32 }
88 if c_y < 0 { nx_genver_emit("missing_y" as *u8, 1); return 33 }
89
90 let d: i64 = ne_x[0]
91 let n_tok: i64 = ne_x[1]
92 if ne_w[0] != d { nx_genver_emit("norm_w_dim_mismatch" as *u8, ne_w[0]); return 34 }
93 if ne_y[0] != d { nx_genver_emit("out_dim_mismatch" as *u8, ne_y[0]); return 35 }
94 // A wrong chunk count still yields plausible numbers, so catch it structurally.
95 if c_a != n_chunks * d {
96 nx_genver_emit("adaln_width_not_nchunks_times_d" as *u8, c_a)
97 nx_genver_emit("expected" as *u8, n_chunks * d)
98 return 36
99 }
100
101 let x: *u8 = nx_genfix_load(model, xn, lx, c_x)
102 if (x as i64) == 0 { nx_genver_emit("load_failed_x" as *u8, 1); return 40 }
103 let w: *u8 = nx_genfix_load(model, wn, lw, c_w)
104 if (w as i64) == 0 { nx_genver_emit("load_failed_norm_w" as *u8, 1); return 41 }
105 let a: *u8 = nx_genfix_load(model, an, la, c_a)
106 if (a as i64) == 0 { nx_genver_emit("load_failed_adaln" as *u8, 1); return 42 }
107 let y: *u8 = nx_genfix_load(model, yn, ly, c_y)
108 if (y as i64) == 0 { nx_genver_emit("load_failed_y" as *u8, 1); return 43 }
109
110 if n_tok < rows { rows = n_tok }
111 nx_genver_emit("dim" as *u8, d)
112 nx_genver_emit("tokens_total" as *u8, n_tok)
113 nx_genver_emit("tokens_checked" as *u8, rows)
114 nx_genver_emit("scale_chunk" as *u8, scale_idx)
115
116 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(eps_recip))
117 let dinv: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(d))
118 let one: i64 = nx_i32_to_f32(1)
119
120 // (1 + scale) depends only on the channel -- hoist it out of the token loop.
121 let sbase: i64 = scale_idx * d
122 let mods: *i64 = sys_mmap(d * 8) as *i64
123 var g: i64 = 0
124 while g < d {
125 mods[g] = __f32_add(one, nx_le_read_u32(a, (sbase + g) * 4))
126 g = g + 1
127 }
128
129 let tol: *i64 = sys_mmap(64) as *i64
130 nx_genver_tols(tol)
131 let c: *i64 = sys_mmap(128) as *i64
132 nx_genver_init(c, 4)
133
134 var t: i64 = 0
135 while t < rows {
136 let base: i64 = t * d
137 var ss: i64 = 0
138 var i: i64 = 0
139 while i < d {
140 let v: i64 = nx_le_read_u32(x, (base + i) * 4)
141 ss = __f32_add(ss, __f32_mul(v, v))
142 i = i + 1
143 }
144 let rms: i64 = nx_f32_sqrt(__f32_add(__f32_mul(ss, dinv), eps))
145 let inv: i64 = nx_f32_div(nx_i32_to_f32(1), rms)
146
147 i = 0
148 while i < d {
149 let flat: i64 = base + i
150 let v: i64 = nx_le_read_u32(x, flat * 4)
151 let wi: i64 = nx_le_read_u32(w, i * 4)
152 let normed: i64 = __f32_mul(__f32_mul(v, inv), wi)
153 let got: i64 = __f32_mul(normed, mods[i])
154 nx_genver_tally(c, tol, got, nx_le_read_u32(y, flat * 4), flat)
155 i = i + 1
156 }
157 t = t + 1
158 }
159
160 return nx_genver_report(c)
161}