code wiki / (root) / nx_gen_prenorm_mod_verify.nx

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}