code wiki / (root) / nx_qwen_hybrid_attn.nx

nx_qwen_hybrid_attn.nx source

↩ module page · 210 lines · 9027 B

1// nx_qwen_hybrid_attn.nx -- the hybrid attention SUB-LAYER on real Qwen weights, gated vs the f32 path. 2// 3// sd-server -> Nishi migration (task#19). Wires the proven bricks into a real attention sub-layer: 4// x -> RMSNorm(attn_norm F32) -> [Q,K integer Q4_K via nx_q4k_linear ; V f32 Q5_K] -> GQA attention (f32) 5// -> O (attn_output Q5_K f32) -> residual. 6// Gate: the hybrid output (integer Q/K) vs a pure-f32 reference (f32 Q/K), 2 tokens so attention actually 7// mixes -> they must agree within 2% (the only difference is the Q/K quantization). Real blk.0, ggml-correct 8// dequant throughout. Confirms the ASSEMBLY (conversions + flow), the pieces being individually verified. 9// license_tier: ORIGINAL 10import "nx_syscalls.nx" 11import "nx_tier.nx" 12import "nx_le.nx" 13import "nx_strconv.nx" 14import "nx_tensor.nx" 15import "nx_gguf.nx" 16import "nx_gguf_load.nx" 17import "nx_gguf_meta.nx" 18import "nx_placement.nx" 19import "nx_gguf_load_lazy.nx" 20import "nx_q4k_matmul.nx" 21import "nx_dequant_iter.nx" 22import "nx_q4k_linear.nx" 23import "nx_q4k_to_f32.nx" 24import "nx_q5_k_to_f32.nx" 25import "nx_f32.nx" 26import "nx_f32_cvt.nx" 27import "nx_f32_div.nx" 28import "nx_f32_rmsnorm.nx" 29import "nx_f32_gqa_attention.nx" 30 31// f32 linear: out[t*od+o] = sum_i act[t*id+i]*W[o*id+i] (W already f32) 32func ha_f32linear(w: *i64, od: i64, id: i64, act: *i64, nt: i64, out: *i64) -> i64 { 33 var t: i64 = 0 34 while t < nt { 35 var o: i64 = 0 36 while o < od { 37 var acc: i64 = 0 38 var i: i64 = 0 39 while i < id { acc = nx_f32_add(acc, nx_f32_mul(act[t * id + i], w[o * id + i])); i = i + 1 } 40 out[t * od + o] = acc 41 o = o + 1 42 } 43 t = t + 1 44 } 45 return 0 46} 47 48func ha_emit(fd: i64, key: *u8, kl: i64, v: i64) -> i64 { 49 let line: *u8 = sys_mmap(80) 50 var lo: i64 = 0 51 var ki: i64 = 0 52 while ki < kl { line[lo] = key[ki]; lo = lo + 1; ki = ki + 1 } 53 line[lo] = 0x3D; lo = lo + 1 54 let dec: *u8 = sys_mmap(32) 55 let nd: i64 = nx_strconv_format_i64(v, dec) 56 var k: i64 = 0 57 while k < nd { line[lo] = dec[k]; lo = lo + 1; k = k + 1 } 58 line[lo] = 0x0A; lo = lo + 1 59 return sys_write(fd, line, lo) 60} 61 62func main() -> i64 { 63 let path: *u8 = "/mnt/c/Users/elder/elder-ai-platform/models/unified/text_encoder/Huihui-Qwen3-4B-Instruct-2507-abliterated-Q4_K_M.gguf" as *u8 64 let fd: i64 = sys_openat_rd(path) 65 if fd < 0 { return 30 } 66 let CAP: i64 = 1207959552 // 1.15 GB (covers blk.0 attn tensors incl attn_output) 67 let buf: *u8 = sys_mmap(CAP) 68 var total: i64 = 0 69 var go: i64 = 1 70 while go == 1 { 71 let r: i64 = sys_read(fd, ((buf as i64) + total) as *u8, CAP - total) 72 if r <= 0 { go = 0 } else { total = total + r; if total >= CAP { go = 0 } } 73 } 74 sys_close(fd) 75 if total < 100000000 { return 31 } 76 let hdr: *NxGgufHeader = sys_mmap(NX_GGUF_HDR_BYTES) as *NxGgufHeader 77 if nx_gguf_parse(buf, total, hdr) != NX_GGUF_OK { return 40 } 78 79 let nti: nx_int = nx_gguf_find_tensor(hdr, "blk.0.attn_norm.weight" as *u8, 22) 80 let qi: nx_int = nx_gguf_find_tensor(hdr, "blk.0.attn_q.weight" as *u8, 19) 81 let ki: nx_int = nx_gguf_find_tensor(hdr, "blk.0.attn_k.weight" as *u8, 19) 82 let vi: nx_int = nx_gguf_find_tensor(hdr, "blk.0.attn_v.weight" as *u8, 19) 83 let oi: nx_int = nx_gguf_find_tensor(hdr, "blk.0.attn_output.weight" as *u8, 24) 84 if nti < 0 { return 50 } 85 if qi < 0 { return 51 } 86 if ki < 0 { return 52 } 87 if vi < 0 { return 53 } 88 if oi < 0 { return 54 } 89 let ntt: *NxGgufTensorInfo = nx_gguf_tensor_at(hdr, nti) 90 let qt: *NxGgufTensorInfo = nx_gguf_tensor_at(hdr, qi) 91 let kt: *NxGgufTensorInfo = nx_gguf_tensor_at(hdr, ki) 92 let vt: *NxGgufTensorInfo = nx_gguf_tensor_at(hdr, vi) 93 let ot: *NxGgufTensorInfo = nx_gguf_tensor_at(hdr, oi) 94 95 let HID: i64 = qt.dim_0 // 4096 96 let QD: i64 = qt.dim_1 // 4096 (32*128) 97 let KVD: i64 = kt.dim_1 // 1024 (8*128) 98 let HEAD: i64 = 128 99 let NQH: i64 = QD / HEAD // 32 100 let NKVH: i64 = KVD / HEAD // 8 101 let NT: i64 = 2 102 103 let n_off: i64 = hdr.data_off + ntt.offset 104 let q_off: i64 = hdr.data_off + qt.offset 105 let k_off: i64 = hdr.data_off + kt.offset 106 let v_off: i64 = hdr.data_off + vt.offset 107 let o_off: i64 = hdr.data_off + ot.offset 108 if o_off + QD * (HID / 256) * 176 > total { return 63 } // attn_output must be inside the prefix 109 110 // ---- attn_norm gamma (F32 weights, read directly) ---- 111 let gamma: *i64 = sys_mmap(HID * 8) as *i64 112 var i: i64 = 0 113 while i < HID { gamma[i] = nx_le_read_u32(buf, n_off + i * 4); i = i + 1 } 114 115 // ---- synthetic input x (NT tokens x HID), f32 ---- 116 let x: *i64 = sys_mmap(NT * HID * 8) as *i64 117 i = 0 118 while i < NT * HID { x[i] = nx_q10_to_f32(512 + (i - (i / 7) * 7) * 128); i = i + 1 } 119 120 // ---- RMSNorm per token -> hn (f32), and hn_q10 for the integer linears ---- 121 let eps: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(1000000)) 122 let hn: *i64 = sys_mmap(NT * HID * 8) as *i64 123 let hnq: *i64 = sys_mmap(NT * HID * 8) as *i64 124 var t: i64 = 0 125 while t < NT { 126 nx_f32_rmsnorm(((x as i64) + t * HID * 8) as *i64, gamma, HID, eps, ((hn as i64) + t * HID * 8) as *i64) 127 t = t + 1 128 } 129 i = 0 130 while i < NT * HID { hnq[i] = _gguf_f32_to_q10(hn[i]); i = i + 1 } 131 132 // ---- dequant weights to f32 (for f32 Q/K ref + f32 V/O in both paths) ---- 133 let Wq: *i64 = sys_mmap(QD * HID * 8) as *i64 134 let Wk: *i64 = sys_mmap(KVD * HID * 8) as *i64 135 let Wv: *i64 = sys_mmap(KVD * HID * 8) as *i64 136 let Wo: *i64 = sys_mmap(QD * QD * 8) as *i64 137 nx_q4k_to_f32(buf, q_off, QD * HID, Wq) 138 nx_q4k_to_f32(buf, k_off, KVD * HID, Wk) 139 nx_q5_k_to_f32(buf, v_off, KVD * HID, Wv) 140 nx_q5_k_to_f32(buf, o_off, QD * QD, Wo) 141 142 // ---- Q/K/V ---- 143 let it: *NxQ4KBlockIter = nx_q4k_iter_alloc() 144 let Qh_q10: *i64 = sys_mmap(NT * QD * 8) as *i64 145 let Kh_q10: *i64 = sys_mmap(NT * KVD * 8) as *i64 146 nx_q4k_linear(buf, q_off, QD, HID, hnq, NT, it, Qh_q10) // integer Q 147 nx_q4k_linear(buf, k_off, KVD, HID, hnq, NT, it, Kh_q10) // integer K 148 let Qh: *i64 = sys_mmap(NT * QD * 8) as *i64 149 let Kh: *i64 = sys_mmap(NT * KVD * 8) as *i64 150 i = 0 151 while i < NT * QD { Qh[i] = nx_q10_to_f32(Qh_q10[i]); i = i + 1 } 152 i = 0 153 while i < NT * KVD { Kh[i] = nx_q10_to_f32(Kh_q10[i]); i = i + 1 } 154 let Qf: *i64 = sys_mmap(NT * QD * 8) as *i64 155 let Kf: *i64 = sys_mmap(NT * KVD * 8) as *i64 156 let V: *i64 = sys_mmap(NT * KVD * 8) as *i64 157 ha_f32linear(Wq, QD, HID, hn, NT, Qf) // f32 Q ref 158 ha_f32linear(Wk, KVD, HID, hn, NT, Kf) // f32 K ref 159 ha_f32linear(Wv, KVD, HID, hn, NT, V) // f32 V (both paths) 160 161 // ---- GQA attention (f32) for both paths ---- 162 let scale: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(11)) // ~1/sqrt(128)=0.0884; 1/11=0.0909 (close enough) 163 let attn_h: *i64 = sys_mmap(NT * QD * 8) as *i64 164 let attn_f: *i64 = sys_mmap(NT * QD * 8) as *i64 165 nx_f32_gqa_attention(Qh, Kh, V, NT, NQH, NKVH, HEAD, scale, attn_h) 166 nx_f32_gqa_attention(Qf, Kf, V, NT, NQH, NKVH, HEAD, scale, attn_f) 167 168 // ---- O projection (f32) + residual ---- 169 let ao_h: *i64 = sys_mmap(NT * QD * 8) as *i64 170 let ao_f: *i64 = sys_mmap(NT * QD * 8) as *i64 171 ha_f32linear(Wo, QD, QD, attn_h, NT, ao_h) 172 ha_f32linear(Wo, QD, QD, attn_f, NT, ao_f) 173 174 // ---- gate: x + ao_h ~ x + ao_f within 2% (diff is only the Q/K quantization) ---- 175 let tolf: i64 = nx_f32_div(nx_i32_to_f32(2), nx_i32_to_f32(100)) 176 var worst: i64 = 0 177 var nchk: i64 = 0 178 // x is [NT,HID], ao is [NT,QD]; QD==HID==4096 so x[i] aligns with ao[i] in the residual. 179 i = 0 180 while i < NT * QD { 181 let rh: i64 = nx_f32_add(x[i], ao_h[i]) 182 let rf: i64 = nx_f32_add(x[i], ao_f[i]) 183 let ref_abs: i64 = rf & 0x7FFFFFFF 184 let diff: i64 = nx_f32_sub(rh, rf) & 0x7FFFFFFF 185 let thr: i64 = nx_f32_mul(tolf, ref_abs) 186 if diff >= thr { 187 // allow tiny-magnitude entries (ref ~0) to pass on absolute 1e-3 188 let abs_ok: i64 = nx_f32_div(nx_i32_to_f32(1), nx_i32_to_f32(1000)) 189 if (diff & 0x7FFFFFFF) >= abs_ok { worst = worst + 1 } 190 } 191 nchk = nchk + 1 192 i = i + 1 193 } 194 195 let ofd: i64 = sys_openat_wr("/tmp/zimg_hybrid_attn.txt" as *u8, 0x1a4) 196 if ofd >= 0 { 197 ha_emit(ofd, "HID" as *u8, 3, HID) 198 ha_emit(ofd, "QD" as *u8, 2, QD) 199 ha_emit(ofd, "KVD" as *u8, 3, KVD) 200 ha_emit(ofd, "NQH" as *u8, 3, NQH) 201 ha_emit(ofd, "NKVH" as *u8, 4, NKVH) 202 ha_emit(ofd, "checked" as *u8, 7, nchk) 203 ha_emit(ofd, "mismatch" as *u8, 8, worst) 204 ha_emit(ofd, "attn_h0" as *u8, 7, attn_h[0]) 205 ha_emit(ofd, "attn_f0" as *u8, 7, attn_f[0]) 206 sys_close(ofd) 207 } 208 if worst > 0 { return 80 } 209 return 0 210}