nx_gen_rmsnorm_verify.nx source
↩ module page · 132 lines · 5225 B
1// nx_gen_rmsnorm_verify.nx -- SOVEREIGN RMSNorm, verified vs the oracle.
2//
3// out[t][i] = x[t][i] / sqrt(mean_i(x[t]^2) + eps) * w[i]
4//
5// RMSNorm appears four times per DiT block and in every Llama-class text encoder, so this is one
6// of the highest-multiplicity ops in the whole gen stack. Named for the OP; the model, the tap
7// names and eps are all arguments.
8//
9// Usage: nx_gen_rmsnorm_verify <model_id> <x_name> <w_name> <y_name> [rows] [eps_recip]
10// eps_recip eps expressed as 1/N (default 1000000 = 1e-6, the ggml RMSNorm default). Passed as
11// an integer reciprocal because eps is always a negative power of ten in practice and
12// an argv float parser is a defect surface this organ does not need.
13//
14// ONE THING TO KNOW ABOUT THE ORACLE FIXTURE: in this engine RMSNorm ends with ggml_mul_inplace,
15// so its result is a VIEW of the rms_norm output. Tapping a view without protecting its storage
16// dumps recycled memory -- see reference-inplace-view-tap-defect-2026-08-06. The tap now flags the
17// whole view chain, which is why these fixtures are trustworthy at all.
18//
19// Hot loop uses the HARDWARE __f32_* intrinsics, not the nx_f32_* software IEEE-754 twins: the
20// software pair measured 9-10x slower in this lane's matmul and is the single easiest way to make
21// a correct sovereign kernel look like a failed one.
22// license_tier: ORIGINAL
23
24import "nx_syscalls.nx"
25import "nx_le.nx"
26import "nx_f32.nx"
27import "nx_f32_div.nx"
28import "nx_f32_cvt.nx"
29import "nx_strconv.nx"
30import "nx_genfix.nx"
31import "nx_genver.nx"
32const K_MAGIC_1000000: i64 = 1000000
33
34func zr_strlen(s: *u8) -> i64 {
35 var n: i64 = 0
36 while s[n] != (0 as u8) { n = n + 1 }
37 return n
38}
39
40func main(argc: i64, argv: *i64) -> i64 {
41 if argc < 5 {
42 nx_genver_emit("usage_model_x_w_y_rows_epsrecip" as *u8, argc)
43 return 2
44 }
45 let model: *u8 = argv[1] as *u8
46 let xn: *u8 = argv[2] as *u8
47 let wn: *u8 = argv[3] as *u8
48 let yn: *u8 = argv[4] as *u8
49 let lx: i64 = zr_strlen(xn)
50 let lw: i64 = zr_strlen(wn)
51 let ly: i64 = zr_strlen(yn)
52
53 let errp: *i64 = sys_mmap(32) as *i64
54 var rows: i64 = 4
55 if argc >= 6 {
56 errp[0] = 0
57 rows = nx_strconv_parse_i64(argv[5] as *u8, errp)
58 if errp[0] != 0 { nx_genver_emit("bad_rows" as *u8, 1); return 3 }
59 }
60 var eps_recip: i64 = K_MAGIC_1000000
61 if argc >= 7 {
62 errp[0] = 0
63 eps_recip = nx_strconv_parse_i64(argv[6] as *u8, errp)
64 if errp[0] != 0 { nx_genver_emit("bad_eps_recip" as *u8, 1); return 4 }
65 if eps_recip <= 0 { nx_genver_emit("bad_eps_recip" as *u8, eps_recip); return 5 }
66 }
67
68 let ne_x: *i64 = sys_mmap(64) as *i64
69 let ne_w: *i64 = sys_mmap(64) as *i64
70 let ne_y: *i64 = sys_mmap(64) as *i64
71 let c_x: i64 = nx_genfix_dims(model, xn, lx, ne_x)
72 let c_w: i64 = nx_genfix_dims(model, wn, lw, ne_w)
73 let c_y: i64 = nx_genfix_dims(model, yn, ly, ne_y)
74 if c_x < 0 { nx_genver_emit("missing_x" as *u8, 1); return 30 }
75 if c_w < 0 { nx_genver_emit("missing_w" as *u8, 1); return 31 }
76 if c_y < 0 { nx_genver_emit("missing_y" as *u8, 1); return 32 }
77
78 let d: i64 = ne_x[0]
79 let n_tok: i64 = ne_x[1]
80 if ne_w[0] != d { nx_genver_emit("weight_dim_mismatch" as *u8, ne_w[0]); return 33 }
81 if ne_y[0] != d { nx_genver_emit("out_dim_mismatch" as *u8, ne_y[0]); return 34 }
82 if ne_y[1] != n_tok { nx_genver_emit("out_tokens_mismatch" as *u8, ne_y[1]); return 35 }
83
84 let x: *u8 = nx_genfix_load(model, xn, lx, c_x)
85 if (x as i64) == 0 { nx_genver_emit("load_failed_x" as *u8, 1); return 40 }
86 let w: *u8 = nx_genfix_load(model, wn, lw, c_w)
87 if (w as i64) == 0 { nx_genver_emit("load_failed_w" as *u8, 1); return 41 }
88 let y: *u8 = nx_genfix_load(model, yn, ly, c_y)
89 if (y as i64) == 0 { nx_genver_emit("load_failed_y" as *u8, 1); return 42 }
90
91 if n_tok < rows { rows = n_tok }
92 nx_genver_emit("dim" as *u8, d)
93 nx_genver_emit("tokens_total" as *u8, n_tok)
94 nx_genver_emit("tokens_checked" as *u8, rows)
95 nx_genver_emit("eps_recip" as *u8, eps_recip)
96
97 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(eps_recip))
98 let dinv: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(d))
99
100 let tol: *i64 = sys_mmap(64) as *i64
101 nx_genver_tols(tol)
102 let c: *i64 = sys_mmap(128) as *i64
103 nx_genver_init(c, 4)
104
105 var t: i64 = 0
106 while t < rows {
107 let base: i64 = t * d
108 // mean of squares
109 var ss: i64 = 0
110 var i: i64 = 0
111 while i < d {
112 let v: i64 = nx_le_read_u32(x, (base + i) * 4)
113 ss = __f32_add(ss, __f32_mul(v, v))
114 i = i + 1
115 }
116 let rms: i64 = nx_f32_sqrt(__f32_add(__f32_mul(ss, dinv), eps))
117 let inv: i64 = nx_f32_div(nx_i32_to_f32(1), rms)
118
119 i = 0
120 while i < d {
121 let flat: i64 = base + i
122 let v: i64 = nx_le_read_u32(x, flat * 4)
123 let wi: i64 = nx_le_read_u32(w, i * 4)
124 let got: i64 = __f32_mul(__f32_mul(v, inv), wi)
125 nx_genver_tally(c, tol, got, nx_le_read_u32(y, flat * 4), flat)
126 i = i + 1
127 }
128 t = t + 1
129 }
130
131 return nx_genver_report(c)
132}