code wiki / (root) / nx_gen_rope_verify.nx

nx_gen_rope_verify.nx source

↩ module page · 132 lines · 5744 B

1// nx_gen_rope_verify.nx -- SOVEREIGN interleaved RoPE, verified vs the oracle. 2// 3// y[h][l][2j+r] = x[l][h][2j] * pe[l][j][r][0] + x[l][h][2j+1] * pe[l][j][r][1] 4// 5// pe holds a 2x2 rotation per (token l, pair j): [[cos, -sin], [sin, cos]], so r=0 gives 6// x0*cos - x1*sin and r=1 gives x0*sin + x1*cos -- ordinary interleaved rotary embedding. 7// Derived from rope.hpp's apply_rope (permute/reshape/repeat chain), not assumed. 8// 9// Usage: nx_gen_rope_verify <model> <x> <pe> <y> <head_dim> <n_heads> [rows] 10// 11// ⚠THE AXIS PERMUTATION IS THE WHOLE DIFFICULTY, and it is silent when wrong: 12// input x is [head_dim, n_heads, L] -> index (l*n_heads + h)*head_dim + d 13// output y is [head_dim, L, n_heads] -> index (h*L + l)*head_dim + d 14// apply_rope permutes head and token axes. Reading either side in the other's order produces 15// finite, plausible, entirely wrong numbers -- no NaN, no crash, just a silently different 16// tensor. Both strides are therefore written out explicitly here rather than shared. 17// 18// pe memory index: ((l*(head_dim/2) + j)*2 + r)*2 + c -- ne = [2, 2, head_dim/2, L] 19// 20// Hardware __f32_* intrinsics in the hot loop, never the nx_f32_* software twins. 21// license_tier: ORIGINAL 22 23import "nx_syscalls.nx" 24import "nx_le.nx" 25import "nx_f32.nx" 26import "nx_f32_div.nx" 27import "nx_f32_cvt.nx" 28import "nx_strconv.nx" 29import "nx_genfix.nx" 30import "nx_genver.nx" 31 32func zrp_strlen(s: *u8) -> i64 { 33 var n: i64 = 0 34 while s[n] != (0 as u8) { n = n + 1 } 35 return n 36} 37 38func main(argc: i64, argv: *i64) -> i64 { 39 if argc < 7 { 40 nx_genver_emit("usage_model_x_pe_y_headdim_nheads" as *u8, argc) 41 return 2 42 } 43 let model: *u8 = argv[1] as *u8 44 let xn: *u8 = argv[2] as *u8 45 let pn: *u8 = argv[3] as *u8 46 let yn: *u8 = argv[4] as *u8 47 let errp: *i64 = sys_mmap(32) as *i64 48 errp[0] = 0 49 let head_dim: i64 = nx_strconv_parse_i64(argv[5] as *u8, errp) 50 if errp[0] != 0 { nx_genver_emit("bad_head_dim" as *u8, 1); return 3 } 51 errp[0] = 0 52 let n_heads: i64 = nx_strconv_parse_i64(argv[6] as *u8, errp) 53 if errp[0] != 0 { nx_genver_emit("bad_n_heads" as *u8, 1); return 4 } 54 var rows: i64 = 4 55 if argc >= 8 { 56 errp[0] = 0 57 rows = nx_strconv_parse_i64(argv[7] as *u8, errp) 58 if errp[0] != 0 { nx_genver_emit("bad_rows" as *u8, 1); return 5 } 59 } 60 61 let lx: i64 = zrp_strlen(xn) 62 let lp: i64 = zrp_strlen(pn) 63 let ly: i64 = zrp_strlen(yn) 64 let ne_x: *i64 = sys_mmap(64) as *i64 65 let ne_p: *i64 = sys_mmap(64) as *i64 66 let ne_y: *i64 = sys_mmap(64) as *i64 67 let c_x: i64 = nx_genfix_dims(model, xn, lx, ne_x) 68 let c_p: i64 = nx_genfix_dims(model, pn, lp, ne_p) 69 let c_y: i64 = nx_genfix_dims(model, yn, ly, ne_y) 70 if c_x < 0 { nx_genver_emit("missing_x" as *u8, 1); return 30 } 71 if c_p < 0 { nx_genver_emit("missing_pe" as *u8, 1); return 31 } 72 if c_y < 0 { nx_genver_emit("missing_y" as *u8, 1); return 32 } 73 74 let half: i64 = head_dim / 2 75 // x is [head_dim, n_heads, L]; y is [head_dim, L, n_heads]. Assert BOTH, because the whole 76 // point of this organ is that the two differ and a swap is invisible in the numbers. 77 if ne_x[0] != head_dim { nx_genver_emit("x_d0_not_head_dim" as *u8, ne_x[0]); return 33 } 78 if ne_x[1] != n_heads { nx_genver_emit("x_d1_not_n_heads" as *u8, ne_x[1]); return 34 } 79 if ne_y[0] != head_dim { nx_genver_emit("y_d0_not_head_dim" as *u8, ne_y[0]); return 35 } 80 if ne_y[2] != n_heads { nx_genver_emit("y_d2_not_n_heads" as *u8, ne_y[2]); return 36 } 81 let n_tok: i64 = ne_x[2] 82 if ne_y[1] != n_tok { nx_genver_emit("y_d1_not_tokens" as *u8, ne_y[1]); return 37 } 83 if ne_p[2] != half { nx_genver_emit("pe_d2_not_half_head_dim" as *u8, ne_p[2]); return 38 } 84 if ne_p[3] != n_tok { nx_genver_emit("pe_d3_not_tokens" as *u8, ne_p[3]); return 39 } 85 86 let x: *u8 = nx_genfix_load(model, xn, lx, c_x) 87 if (x as i64) == 0 { nx_genver_emit("load_failed_x" as *u8, 1); return 40 } 88 let p: *u8 = nx_genfix_load(model, pn, lp, c_p) 89 if (p as i64) == 0 { nx_genver_emit("load_failed_pe" as *u8, 1); return 41 } 90 let y: *u8 = nx_genfix_load(model, yn, ly, c_y) 91 if (y as i64) == 0 { nx_genver_emit("load_failed_y" as *u8, 1); return 42 } 92 93 if n_tok < rows { rows = n_tok } 94 nx_genver_emit("head_dim" as *u8, head_dim) 95 nx_genver_emit("n_heads" as *u8, n_heads) 96 nx_genver_emit("tokens_total" as *u8, n_tok) 97 nx_genver_emit("tokens_checked" as *u8, rows) 98 99 let tol: *i64 = sys_mmap(64) as *i64 100 nx_genver_tols(tol) 101 let c: *i64 = sys_mmap(128) as *i64 102 nx_genver_init(c, 4) 103 104 var l: i64 = 0 105 while l < rows { 106 var h: i64 = 0 107 while h < n_heads { 108 let xrow: i64 = (l * n_heads + h) * head_dim // [head_dim, n_heads, L] 109 let yrow: i64 = (h * n_tok + l) * head_dim // [head_dim, L, n_heads] 110 var j: i64 = 0 111 while j < half { 112 let x0: i64 = nx_le_read_u32(x, (xrow + 2 * j) * 4) 113 let x1: i64 = nx_le_read_u32(x, (xrow + 2 * j + 1) * 4) 114 var r: i64 = 0 115 while r < 2 { 116 let pbase: i64 = ((l * half + j) * 2 + r) * 2 117 let p0: i64 = nx_le_read_u32(p, pbase * 4) 118 let p1: i64 = nx_le_read_u32(p, (pbase + 1) * 4) 119 let got: i64 = __f32_add(__f32_mul(x0, p0), __f32_mul(x1, p1)) 120 let flat: i64 = yrow + 2 * j + r 121 nx_genver_tally(c, tol, got, nx_le_read_u32(y, flat * 4), flat) 122 r = r + 1 123 } 124 j = j + 1 125 } 126 h = h + 1 127 } 128 l = l + 1 129 } 130 131 return nx_genver_report(c) 132}