nx_gen_qknorm_verify.nx source
↩ module page · 148 lines · 6040 B
1// nx_gen_qknorm_verify.nx -- SOVEREIGN per-head QK-RMSNorm over a PACKED qkv, verified vs oracle.
2//
3// out[t][h][d] = qkv[t][(head_base+h)*head_dim + d] / sqrt(mean_d(.^2) + eps) * w[d]
4//
5// QK-norm normalizes each attention head's query (or key) vector independently over head_dim,
6// reading from the fused qkv projection where q, k and v are packed head-major in one row.
7// Z-Image, Qwen-Image and most recent DiTs do this; it is the first stage of the attention core.
8//
9// Usage:
10// nx_gen_qknorm_verify <model> <qkv> <norm_w> <out> <head_dim> <n_heads> <head_base> [rows] [eps_recip]
11// head_base 0 for q; n_heads for k (q, k, v are packed in that order)
12//
13// WHY THIS IS A SEPARATE ORGAN AND NOT THE RMSNORM ONE: the two tensors have DIFFERENT row
14// strides. The packed qkv row is (n_q + n_k + n_v) * head_dim wide, while the output row is
15// n_heads * head_dim. A generic normalizer that assumed one stride would read the wrong head and
16// still produce finite, plausible numbers -- the failure mode this lane keeps meeting.
17//
18// Hardware __f32_* intrinsics in the hot loop, never the nx_f32_* software twins.
19// license_tier: ORIGINAL
20
21import "nx_syscalls.nx"
22import "nx_le.nx"
23import "nx_f32.nx"
24import "nx_f32_div.nx"
25import "nx_f32_cvt.nx"
26import "nx_strconv.nx"
27import "nx_genfix.nx"
28import "nx_genver.nx"
29const K_MAGIC_1000000: i64 = 1000000
30
31func zq_strlen(s: *u8) -> i64 {
32 var n: i64 = 0
33 while s[n] != (0 as u8) { n = n + 1 }
34 return n
35}
36
37func main(argc: i64, argv: *i64) -> i64 {
38 if argc < 8 {
39 nx_genver_emit("usage_model_qkv_w_out_headdim_nheads_headbase" as *u8, argc)
40 return 2
41 }
42 let model: *u8 = argv[1] as *u8
43 let qn: *u8 = argv[2] as *u8
44 let wn: *u8 = argv[3] as *u8
45 let on: *u8 = argv[4] as *u8
46 let errp: *i64 = sys_mmap(32) as *i64
47 errp[0] = 0
48 let head_dim: i64 = nx_strconv_parse_i64(argv[5] as *u8, errp)
49 if errp[0] != 0 { nx_genver_emit("bad_head_dim" as *u8, 1); return 3 }
50 errp[0] = 0
51 let n_heads: i64 = nx_strconv_parse_i64(argv[6] as *u8, errp)
52 if errp[0] != 0 { nx_genver_emit("bad_n_heads" as *u8, 1); return 4 }
53 errp[0] = 0
54 let head_base: i64 = nx_strconv_parse_i64(argv[7] as *u8, errp)
55 if errp[0] != 0 { nx_genver_emit("bad_head_base" as *u8, 1); return 5 }
56
57 var rows: i64 = 4
58 if argc >= 9 {
59 errp[0] = 0
60 rows = nx_strconv_parse_i64(argv[8] as *u8, errp)
61 if errp[0] != 0 { nx_genver_emit("bad_rows" as *u8, 1); return 6 }
62 }
63 var eps_recip: i64 = K_MAGIC_1000000
64 if argc >= 10 {
65 errp[0] = 0
66 eps_recip = nx_strconv_parse_i64(argv[9] as *u8, errp)
67 if errp[0] != 0 { nx_genver_emit("bad_eps_recip" as *u8, 1); return 7 }
68 }
69
70 let lq: i64 = zq_strlen(qn)
71 let lw: i64 = zq_strlen(wn)
72 let lo: i64 = zq_strlen(on)
73 let ne_q: *i64 = sys_mmap(64) as *i64
74 let ne_w: *i64 = sys_mmap(64) as *i64
75 let ne_o: *i64 = sys_mmap(64) as *i64
76 let c_q: i64 = nx_genfix_dims(model, qn, lq, ne_q)
77 let c_w: i64 = nx_genfix_dims(model, wn, lw, ne_w)
78 let c_o: i64 = nx_genfix_dims(model, on, lo, ne_o)
79 if c_q < 0 { nx_genver_emit("missing_qkv" as *u8, 1); return 30 }
80 if c_w < 0 { nx_genver_emit("missing_norm_w" as *u8, 1); return 31 }
81 if c_o < 0 { nx_genver_emit("missing_out" as *u8, 1); return 32 }
82
83 let qkv_w: i64 = ne_q[0] // packed width, e.g. (30+30+30)*128
84 let n_tok: i64 = ne_q[1]
85 if ne_w[0] != head_dim { nx_genver_emit("norm_w_not_head_dim" as *u8, ne_w[0]); return 33 }
86 // out is [head_dim, n_heads, n_tok]; check the declared geometry against the fixture rather
87 // than trusting the caller's head count.
88 if ne_o[0] != head_dim { nx_genver_emit("out_d0_not_head_dim" as *u8, ne_o[0]); return 34 }
89 if ne_o[1] != n_heads { nx_genver_emit("out_d1_not_n_heads" as *u8, ne_o[1]); return 35 }
90 if (head_base + n_heads) * head_dim > qkv_w {
91 nx_genver_emit("head_window_exceeds_qkv_width" as *u8, qkv_w)
92 return 36
93 }
94
95 let q: *u8 = nx_genfix_load(model, qn, lq, c_q)
96 if (q as i64) == 0 { nx_genver_emit("load_failed_qkv" as *u8, 1); return 40 }
97 let w: *u8 = nx_genfix_load(model, wn, lw, c_w)
98 if (w as i64) == 0 { nx_genver_emit("load_failed_norm_w" as *u8, 1); return 41 }
99 let o: *u8 = nx_genfix_load(model, on, lo, c_o)
100 if (o as i64) == 0 { nx_genver_emit("load_failed_out" as *u8, 1); return 42 }
101
102 if n_tok < rows { rows = n_tok }
103 nx_genver_emit("head_dim" as *u8, head_dim)
104 nx_genver_emit("n_heads" as *u8, n_heads)
105 nx_genver_emit("head_base" as *u8, head_base)
106 nx_genver_emit("qkv_width" as *u8, qkv_w)
107 nx_genver_emit("tokens_checked" as *u8, rows)
108
109 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(eps_recip))
110 let hdinv: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(head_dim))
111
112 let tol: *i64 = sys_mmap(64) as *i64
113 nx_genver_tols(tol)
114 let c: *i64 = sys_mmap(128) as *i64
115 nx_genver_init(c, 4)
116
117 var t: i64 = 0
118 while t < rows {
119 var h: i64 = 0
120 while h < n_heads {
121 let src: i64 = t * qkv_w + (head_base + h) * head_dim
122 let dst: i64 = t * n_heads * head_dim + h * head_dim
123
124 var ss: i64 = 0
125 var d: i64 = 0
126 while d < head_dim {
127 let v: i64 = nx_le_read_u32(q, (src + d) * 4)
128 ss = __f32_add(ss, __f32_mul(v, v))
129 d = d + 1
130 }
131 let rms: i64 = nx_f32_sqrt(__f32_add(__f32_mul(ss, hdinv), eps))
132 let inv: i64 = nx_f32_div(nx_i32_to_f32(1), rms)
133
134 d = 0
135 while d < head_dim {
136 let v: i64 = nx_le_read_u32(q, (src + d) * 4)
137 let wi: i64 = nx_le_read_u32(w, d * 4)
138 let got: i64 = __f32_mul(__f32_mul(v, inv), wi)
139 nx_genver_tally(c, tol, got, nx_le_read_u32(o, (dst + d) * 4), dst + d)
140 d = d + 1
141 }
142 h = h + 1
143 }
144 t = t + 1
145 }
146
147 return nx_genver_report(c)
148}