code wiki / (root) / nx_gen_qknorm_verify.nx

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}