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}