1use super::*;
4
5#[repr(C)]
6#[derive(Clone, Copy)]
7struct AttnGemmParams {
8 kv_row: u32,
9 blocks_per_row: u32,
10 head: u32,
11 kl: u32,
12 kl_pad: u32,
13 scale_off: u32,
14 base: u32,
15 nb: u32,
16 n_rep: u32,
17 scale: f32,
18}
19
20#[repr(C)]
21struct GemmHParams {
22 m: u32,
23 n: u32,
24 k: u32,
25 lda: u32,
26 ldb: u32,
27 ldc: u32,
28}
29
30#[repr(C)]
31struct SelMaskParams {
32 row0: u32,
33 mask_words: u32,
34 ratio: u32,
35 k: u32,
36}
37
38impl Gpu<'_> {
39 pub(super) fn pf_attention(
43 &self,
44 enc: &Enc,
45 a: &Attn,
46 base: usize,
47 t: usize,
48 pf: &PrefillScratch,
49 ) {
50 let c = &self.p.cfg;
51 let hd = c.head_dim;
52 let n_heads = c.num_attention_heads;
53 let n_kv = c.num_key_value_heads;
54 let n_rep = n_heads / n_kv;
55 let kv_row = n_kv * hd;
56 let nbu = t as u32;
57
58 self.qmm(enc, &a.q, &pf.mixed, &pf.qg, t);
59 self.qmm(enc, &a.k, &pf.mixed, &pf.k, t);
60 self.qmm(enc, &a.v, &pf.mixed, &pf.v, t);
61 self.qmm(enc, &a.iqk, &pf.mixed, &pf.iqk, t);
62
63 let rot = (hd as f64 * c.partial_rotary_factor) as u32;
64 let qp = QkRopeParams {
65 n_heads: n_heads as u32,
66 head_dim: hd as u32,
67 stride: 2 * hd as u32,
68 rot,
69 pos: base as u32,
70 theta: c.rope_parameters.rope_theta as f32,
71 eps: c.rms_norm_eps as f32,
72 };
73
74 self.dispatch(
75 enc,
76 &self.pipes.qk_norm_rope_b,
77 |e| {
78 self.bind(e, 0, &pf.qg, 0);
79 self.bind(e, 1, &self.dense, a.qn.0);
80 set_bytes(e, 2, &qp);
81 set_bytes(e, 3, &nbu);
82 },
83 (t * n_heads).div_ceil(4),
84 128,
85 true,
86 );
87
88 let kp = QkRopeParams {
89 n_heads: n_kv as u32,
90 stride: hd as u32,
91 ..qp
92 };
93
94 self.dispatch(
95 enc,
96 &self.pipes.qk_norm_rope_b,
97 |e| {
98 self.bind(e, 0, &pf.k, 0);
99 self.bind(e, 1, &self.dense, a.kn.0);
100 set_bytes(e, 2, &kp);
101 set_bytes(e, 3, &nbu);
102 },
103 (t * n_kv).div_ceil(4),
104 128,
105 true,
106 );
107
108 let ratio = c.indexer_compress_ratio;
111 let ihd = c.indexer_head_dim;
112 let inh = c.indexer_n_heads;
113 let kblk = c.indexer_budget / ratio;
114 let max_blocks = self.max_t / ratio + 1;
115 let mask_words = max_blocks.div_ceil(32);
116 let ip = IndexParams {
117 ihd: ihd as u32,
118 inh: inh as u32,
119 qk_dim: a.iqk.out,
120 ratio: ratio as u32,
121 rot,
122 theta: c.rope_parameters.rope_theta as f32,
123 eps: c.rms_norm_eps as f32,
124 base_pos: base as u32,
125 nb: nbu,
126 b0: (base / ratio) as u32,
127 b1: ((base + t) / ratio) as u32,
128 k: kblk as u32,
129 max_blocks: max_blocks as u32,
130 vis_stride: (c.indexer_budget + ratio) as u32,
131 mask_words: mask_words as u32,
132 };
133
134 self.dispatch(
135 enc,
136 &self.pipes.index_append,
137 |e| {
138 self.bind(e, 0, &pf.iqk, 0);
139 self.bind(e, 1, &a.ikc, 0);
140 set_bytes(e, 2, &ip);
141 },
142 t * ihd,
143 256,
144 false,
145 );
146
147 if ip.b1 > ip.b0 {
148 self.dispatch(
149 enc,
150 &self.pipes.index_blocks,
151 |e| {
152 self.bind(e, 0, &a.ikc, 0);
153 self.bind(e, 1, &a.blk, 0);
154 self.bind(e, 2, &self.dense, a.ikn.0);
155 set_bytes(e, 3, &ip);
156 },
157 (ip.b1 - ip.b0) as usize,
158 ihd,
159 true,
160 );
161 }
162
163 if (base + t) / ratio > kblk {
164 self.dispatch(
165 enc,
166 &self.pipes.index_q,
167 |e| {
168 self.bind(e, 0, &pf.iqk, 0);
169 self.bind(e, 1, &pf.iq, 0);
170 self.bind(e, 2, &self.dense, a.iqn.0);
171 set_bytes(e, 3, &ip);
172 },
173 t * inh,
174 ihd,
175 true,
176 );
177
178 self.dispatch(
179 enc,
180 &self.pipes.index_score,
181 |e| {
182 self.bind(e, 0, &pf.iq, 0);
183 self.bind(e, 1, &a.blk, 0);
184 self.bind(e, 2, &pf.bscore, 0);
185 set_bytes(e, 3, &ip);
186 },
187 t * max_blocks.div_ceil(256),
188 256,
189 true,
190 );
191
192 self.dispatch(
193 enc,
194 &self.pipes.index_select,
195 |e| {
196 self.bind(e, 0, &pf.bscore, 0);
197 self.bind(e, 1, &pf.vis, 0);
198 self.bind(e, 2, &pf.nvis, 0);
199 set_bytes(e, 3, &ip);
200 self.bind(e, 4, &pf.vmask, 0);
201 },
202 t,
203 1024,
204 true,
205 );
206 }
207
208 let scale_off = kv_q8_side(self.max_t, kv_row).0;
209 let kq = KvQParams {
210 row: kv_row as u32,
211 t0: base as u32,
212 nb: nbu,
213 scale_off: scale_off as u32,
214 };
215
216 self.dispatch(
217 enc,
218 &self.pipes.kv_append_q8,
219 |e| {
220 self.bind(e, 0, &pf.k, 0);
221 self.bind(e, 1, &pf.v, 0);
222 self.bind(e, 2, &a.kc, 0);
223 self.bind(e, 3, &a.vc, 0);
224 set_bytes(e, 4, &kq);
225 },
226 (t * (kv_row / 32) * 2).div_ceil(4),
227 128,
228 true,
229 );
230
231 let n_heads_u = n_heads as u32;
233 let mut q0 = 0;
234
235 while q0 < t {
236 let n = (t - q0).min(QS);
237 let base_q = base + q0;
238 let kl = base_q + n;
239 let kl_pad = kl.div_ceil(32) * 32;
240 let p = AttnGemmParams {
241 kv_row: kv_row as u32,
242 blocks_per_row: (kv_row / 32) as u32,
243 head: 0,
244 kl: kl as u32,
245 kl_pad: kl_pad as u32,
246 scale_off: scale_off as u32,
247 base: base_q as u32,
248 nb: n as u32,
249 n_rep: n_rep as u32,
250 scale: (hd as f32).powf(-0.5) * std::f32::consts::LOG2_E,
251 };
252 let qg_off = q0 * n_heads * 2 * hd * 4;
253
254 self.dispatch(
255 enc,
256 &self.pipes.attn_q_stage,
257 |e| {
258 self.bind(e, 0, &pf.qg, qg_off);
259 self.bind(e, 1, &pf.ag_qh, 0);
260 set_bytes(e, 2, &p);
261 set_bytes(e, 3, &n_heads_u);
262 },
263 n_heads * n * hd,
264 256,
265 false,
266 );
267
268 let m = n_rep * n;
269 let sm = SelMaskParams {
270 row0: q0 as u32,
271 mask_words: mask_words as u32,
272 ratio: ratio as u32,
273 k: kblk as u32,
274 };
275
276 for hk in 0..n_kv {
277 let ph = AttnGemmParams {
278 head: hk as u32,
279 ..p
280 };
281
282 self.dispatch(
283 enc,
284 &self.pipes.attn_kv_stage,
285 |e| {
286 self.bind(e, 0, &a.kc, 0);
287 self.bind(e, 1, &a.vc, 0);
288 self.bind(e, 2, &pf.ag_kh, 0);
289 self.bind(e, 3, &pf.ag_vt, 0);
290 set_bytes(e, 4, &ph);
291 },
292 kl_pad * hd,
293 256,
294 false,
295 );
296
297 let ps = GemmHParams {
298 m: m as u32,
299 n: kl_pad as u32,
300 k: hd as u32,
301 lda: hd as u32,
302 ldb: hd as u32,
303 ldc: kl_pad as u32,
304 };
305
306 self.dispatch(
307 enc,
308 &self.pipes.gemm_hh,
309 |e| {
310 self.bind(e, 0, &pf.ag_qh, hk * n_rep * n * hd * 2);
311 self.bind(e, 1, &pf.ag_kh, 0);
312 self.bind(e, 2, &pf.ag_s, 0);
313 set_bytes(e, 3, &ps);
314 },
315 m.div_ceil(128) * (kl_pad / 32),
316 128,
317 true,
318 );
319
320 self.dispatch(
321 enc,
322 &self.pipes.softmax_sel,
323 |e| {
324 self.bind(e, 0, &pf.ag_s, 0);
325 self.bind(e, 1, &pf.ag_p, 0);
326 set_bytes(e, 2, &ph);
327 self.bind(e, 3, &pf.vmask, 0);
328 set_bytes(e, 4, &sm);
329 },
330 m,
331 256,
332 true,
333 );
334
335 let po = GemmHParams {
336 m: m as u32,
337 n: hd as u32,
338 k: kl_pad as u32,
339 lda: kl_pad as u32,
340 ldb: kl_pad as u32,
341 ldc: hd as u32,
342 };
343
344 self.dispatch(
345 enc,
346 &self.pipes.gemm_hh,
347 |e| {
348 self.bind(e, 0, &pf.ag_p, 0);
349 self.bind(e, 1, &pf.ag_vt, 0);
350 self.bind(e, 2, &pf.ag_o, 0);
351 set_bytes(e, 3, &po);
352 },
353 m.div_ceil(128) * (hd / 32),
354 128,
355 true,
356 );
357
358 self.dispatch(
359 enc,
360 &self.pipes.attn_o_scatter,
361 |e| {
362 self.bind(e, 0, &pf.ag_o, 0);
363 self.bind(e, 1, &pf.qg, qg_off);
364 self.bind(e, 2, &pf.attn_out, q0 * n_heads * hd * 4);
365 set_bytes(e, 3, &ph);
366 set_bytes(e, 4, &n_heads_u);
367 },
368 n_rep * n * hd,
369 256,
370 false,
371 );
372 }
373
374 q0 += n;
375 }
376
377 self.qmm(enc, &a.o, &pf.attn_out, &pf.mix_out, t);
378 }
379}