code wiki / (root) / nx_gen_sdpa_verify.nx

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}