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}