Skip to main content

cherenkov/qwen4_exp/gpu/prefill/
attention.rs

1//! QSA masks and GEMM attention over query sub-chunks.
2
3use 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    /// Full attention over t rows at positions base..: GEMM form in
40    /// sub-chunks of QS queries, with the QSA mask for rows past the
41    /// budget. Input pf.mixed, output pf.mix_out.
42    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        // Indexer: raw keys, block keys, and block selection for rows past
109        // the budget (as a bitmask for the softmax).
110        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        // GEMM attention per query sub-chunk and kv head.
232        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}