code wiki / (root) / nx_gen_rmsnorm_verify.nx

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}