nx_gen_sdpa_verify.nx source
↩ module page · 186 lines · 7735 B
1// nx_gen_sdpa_verify.nx -- SOVEREIGN scaled dot-product attention, verified vs the oracle.
2//
3// scores[lk] = sum_d q[h][lq][d] * k[h][lk][d] / sqrt(head_dim)
4// out[lq][h][d] = sum_lk softmax(scores)[lk] * v[lk][h][d]
5//
6// THE SCALE, derived not guessed: ggml_ext_attention_ext uses scale = 1/sqrt(d_head). It also
7// takes a kv_scale (1/128 here) which it applies to k and v before an f16 cast to avoid overflow
8// -- but that cancels EXACTLY: the flash path passes scale/kv_scale into the softmax and then
9// multiplies the result by 1/kv_scale, and the fallback path never applies kv_scale at all. So the
10// arithmetic to reproduce is plain SDPA with 1/sqrt(head_dim), and reading kv_scale as the softmax
11// scale would have been wrong by 11x while still producing a finite, plausible tensor.
12//
13// Usage: nx_gen_sdpa_verify <model> <q> <k> <qkv> <out> <head_dim> <n_heads> <v_head_base> [rows]
14//
15// ⚠THREE DIFFERENT LAYOUTS MEET HERE, which is the whole difficulty:
16// q, k [head_dim, L, n_heads] -> (h*L + l)*head_dim + d (post-RoPE order)
17// v packed qkv [qkv_width, L] -> l*qkv_width + (base+h)*head_dim + d
18// out [n_heads*head_dim, L] -> l*(n_heads*head_dim) + h*head_dim + d
19// RoPE permuted q/k but v is still read straight from the fused projection, so v does NOT share
20// q's layout. Every index is written out explicitly rather than shared.
21//
22// Softmax subtracts the row max before exponentiating -- without it, logits of a few hundred
23// overflow f32 exp and the whole row becomes NaN or zero.
24// license_tier: ORIGINAL
25
26import "nx_syscalls.nx"
27import "nx_le.nx"
28import "nx_f32.nx"
29import "nx_f32_div.nx"
30import "nx_f32_cvt.nx"
31import "nx_f32_exp.nx"
32import "nx_strconv.nx"
33import "nx_genfix.nx"
34import "nx_genver.nx"
35
36func zs_strlen(s: *u8) -> i64 {
37 var n: i64 = 0
38 while s[n] != (0 as u8) { n = n + 1 }
39 return n
40}
41
42func main(argc: i64, argv: *i64) -> i64 {
43 if argc < 9 {
44 nx_genver_emit("usage_model_q_k_qkv_out_headdim_nheads_vbase" as *u8, argc)
45 return 2
46 }
47 let model: *u8 = argv[1] as *u8
48 let qn: *u8 = argv[2] as *u8
49 let kn: *u8 = argv[3] as *u8
50 let vn: *u8 = argv[4] as *u8
51 let on: *u8 = argv[5] as *u8
52 let errp: *i64 = sys_mmap(32) as *i64
53 errp[0] = 0
54 let head_dim: i64 = nx_strconv_parse_i64(argv[6] as *u8, errp)
55 if errp[0] != 0 { nx_genver_emit("bad_head_dim" as *u8, 1); return 3 }
56 errp[0] = 0
57 let n_heads: i64 = nx_strconv_parse_i64(argv[7] as *u8, errp)
58 if errp[0] != 0 { nx_genver_emit("bad_n_heads" as *u8, 1); return 4 }
59 errp[0] = 0
60 let v_base: i64 = nx_strconv_parse_i64(argv[8] as *u8, errp)
61 if errp[0] != 0 { nx_genver_emit("bad_v_head_base" as *u8, 1); return 5 }
62 var rows: i64 = 2
63 if argc >= 10 {
64 errp[0] = 0
65 rows = nx_strconv_parse_i64(argv[9] as *u8, errp)
66 if errp[0] != 0 { nx_genver_emit("bad_rows" as *u8, 1); return 6 }
67 }
68
69 let lq: i64 = zs_strlen(qn)
70 let lk: i64 = zs_strlen(kn)
71 let lv: i64 = zs_strlen(vn)
72 let lo: i64 = zs_strlen(on)
73 let ne_q: *i64 = sys_mmap(64) as *i64
74 let ne_k: *i64 = sys_mmap(64) as *i64
75 let ne_v: *i64 = sys_mmap(64) as *i64
76 let ne_o: *i64 = sys_mmap(64) as *i64
77 let c_q: i64 = nx_genfix_dims(model, qn, lq, ne_q)
78 let c_k: i64 = nx_genfix_dims(model, kn, lk, ne_k)
79 let c_v: i64 = nx_genfix_dims(model, vn, lv, ne_v)
80 let c_o: i64 = nx_genfix_dims(model, on, lo, ne_o)
81 if c_q < 0 { nx_genver_emit("missing_q" as *u8, 1); return 30 }
82 if c_k < 0 { nx_genver_emit("missing_k" as *u8, 1); return 31 }
83 if c_v < 0 { nx_genver_emit("missing_qkv" as *u8, 1); return 32 }
84 if c_o < 0 { nx_genver_emit("missing_out" as *u8, 1); return 33 }
85
86 let n_tok: i64 = ne_q[1]
87 let qkv_w: i64 = ne_v[0]
88 let out_w: i64 = n_heads * head_dim
89 if ne_q[0] != head_dim { nx_genver_emit("q_d0_not_head_dim" as *u8, ne_q[0]); return 34 }
90 if ne_q[2] != n_heads { nx_genver_emit("q_d2_not_n_heads" as *u8, ne_q[2]); return 35 }
91 if ne_o[0] != out_w { nx_genver_emit("out_d0_not_nheads_headdim" as *u8, ne_o[0]); return 36 }
92 if (v_base + n_heads) * head_dim > qkv_w {
93 nx_genver_emit("v_window_exceeds_qkv_width" as *u8, qkv_w)
94 return 37
95 }
96
97 let q: *u8 = nx_genfix_load(model, qn, lq, c_q)
98 if (q as i64) == 0 { nx_genver_emit("load_failed_q" as *u8, 1); return 40 }
99 let k: *u8 = nx_genfix_load(model, kn, lk, c_k)
100 if (k as i64) == 0 { nx_genver_emit("load_failed_k" as *u8, 1); return 41 }
101 let v: *u8 = nx_genfix_load(model, vn, lv, c_v)
102 if (v as i64) == 0 { nx_genver_emit("load_failed_qkv" as *u8, 1); return 42 }
103 let o: *u8 = nx_genfix_load(model, on, lo, c_o)
104 if (o as i64) == 0 { nx_genver_emit("load_failed_out" as *u8, 1); return 43 }
105
106 if n_tok < rows { rows = n_tok }
107 nx_genver_emit("head_dim" as *u8, head_dim)
108 nx_genver_emit("n_heads" as *u8, n_heads)
109 nx_genver_emit("tokens_total" as *u8, n_tok)
110 nx_genver_emit("tokens_checked" as *u8, rows)
111 nx_genver_emit("v_head_base" as *u8, v_base)
112
113 // 1/sqrt(head_dim)
114 let scale: i64 = nx_f32_div(nx_i32_to_f32(1), nx_f32_sqrt(nx_i32_to_f32(head_dim)))
115
116 let tol: *i64 = sys_mmap(64) as *i64
117 nx_genver_tols(tol)
118 let c: *i64 = sys_mmap(128) as *i64
119 nx_genver_init(c, 3)
120
121 let sc: *i64 = sys_mmap(n_tok * 8) as *i64
122 let acc: *i64 = sys_mmap(head_dim * 8) as *i64
123
124 var h: i64 = 0
125 while h < n_heads {
126 var tq: i64 = 0
127 while tq < rows {
128 let qrow: i64 = (h * n_tok + tq) * head_dim
129 // scores over every key position
130 var mx: i64 = 0
131 var tk: i64 = 0
132 while tk < n_tok {
133 let krow: i64 = (h * n_tok + tk) * head_dim
134 var s: i64 = 0
135 var d: i64 = 0
136 while d < head_dim {
137 s = __f32_add(s, __f32_mul(nx_le_read_u32(q, (qrow + d) * 4),
138 nx_le_read_u32(k, (krow + d) * 4)))
139 d = d + 1
140 }
141 s = __f32_mul(s, scale)
142 sc[tk] = s
143 if tk == 0 { mx = s } else { if nx_f32_lt(mx, s) == 1 { mx = s } }
144 tk = tk + 1
145 }
146 // softmax with the row max subtracted
147 // There is no __f32_sub intrinsic (only add/mul/div), so negate by flipping the sign
148 // bit and add -- the idiom nx_q8_0_from_f32 already uses for the same reason.
149 let neg_mx: i64 = mx ^ 0x80000000
150 var sum: i64 = 0
151 tk = 0
152 while tk < n_tok {
153 let e: i64 = nx_f32_exp(__f32_add(sc[tk], neg_mx))
154 sc[tk] = e
155 sum = __f32_add(sum, e)
156 tk = tk + 1
157 }
158 let inv: i64 = nx_f32_div(nx_i32_to_f32(1), sum)
159
160 var d2: i64 = 0
161 while d2 < head_dim { acc[d2] = 0; d2 = d2 + 1 }
162 tk = 0
163 while tk < n_tok {
164 let wgt: i64 = __f32_mul(sc[tk], inv)
165 let vrow: i64 = tk * qkv_w + (v_base + h) * head_dim
166 d2 = 0
167 while d2 < head_dim {
168 acc[d2] = __f32_add(acc[d2], __f32_mul(wgt, nx_le_read_u32(v, (vrow + d2) * 4)))
169 d2 = d2 + 1
170 }
171 tk = tk + 1
172 }
173
174 d2 = 0
175 while d2 < head_dim {
176 let flat: i64 = tq * out_w + h * head_dim + d2
177 nx_genver_tally(c, tol, acc[d2], nx_le_read_u32(o, flat * 4), flat)
178 d2 = d2 + 1
179 }
180 tq = tq + 1
181 }
182 h = h + 1
183 }
184
185 return nx_genver_report(c)
186}