Skip to main content

cherenkov/qwen4_exp/gpu/
attention.rs

1//! QSA indexing, q8 KV append, and selected decode attention.
2
3use super::*;
4
5impl Gpu<'_> {
6    /// Full attention over `nb` rows at positions base_pos..: dense while
7    /// a row's context fits the indexer budget, otherwise over the blocks
8    /// the QSA indexer selects for that row. Input prepped in
9    /// scratch.hc.h1; output in scratch.mix_out.
10    pub(super) fn attention_b(&self, enc: &Enc, a: &Attn, base_pos: usize, nb: usize) {
11        let c = &self.p.cfg;
12        let s = &self.scratch;
13        let hd = c.head_dim as u32;
14        let kv_row = c.num_key_value_heads * c.head_dim;
15        let nbu = nb as u32;
16
17        self.qmv_h(enc, &a.q, &s.qg, nb, &s.hc.h1);
18        self.qmv_h(enc, &a.k, &s.k, nb, &s.hc.h1);
19        self.qmv_h(enc, &a.v, &s.v, nb, &s.hc.h1);
20        self.qmv_h(enc, &a.iqk, &s.iqk, nb, &s.hc.h1);
21
22        let rot = (c.head_dim as f64 * c.partial_rotary_factor) as u32;
23        // Indexer: cache this batch's raw keys, refresh the blocks it
24        // completes, and pick blocks for rows past the budget.
25        let ratio = c.indexer_compress_ratio;
26        let ihd = c.indexer_head_dim;
27        let inh = c.indexer_n_heads;
28        let kblk = c.indexer_budget / ratio;
29        let max_blocks = self.max_t / ratio + 1;
30        let vis_stride = c.indexer_budget + ratio;
31        let ip = IndexParams {
32            ihd: ihd as u32,
33            inh: inh as u32,
34            qk_dim: a.iqk.out,
35            ratio: ratio as u32,
36            rot,
37            theta: c.rope_parameters.rope_theta as f32,
38            eps: c.rms_norm_eps as f32,
39            base_pos: base_pos as u32,
40            nb: nbu,
41            b0: (base_pos / ratio) as u32,
42            b1: ((base_pos + nb) / ratio) as u32,
43            k: kblk as u32,
44            max_blocks: max_blocks as u32,
45            vis_stride: vis_stride as u32,
46            mask_words: max_blocks.div_ceil(32) as u32,
47        };
48
49        self.dispatch(
50            enc,
51            &self.pipes.index_append,
52            |e| {
53                self.bind(e, 0, &s.iqk, 0);
54                self.bind(e, 1, &a.ikc, 0);
55                set_bytes(e, 2, &ip);
56            },
57            nb * ihd,
58            256,
59            false,
60        );
61
62        if ip.b1 > ip.b0 {
63            self.dispatch(
64                enc,
65                &self.pipes.index_blocks,
66                |e| {
67                    self.bind(e, 0, &a.ikc, 0);
68                    self.bind(e, 1, &a.blk, 0);
69                    self.bind(e, 2, &self.dense, a.ikn.0);
70                    set_bytes(e, 3, &ip);
71                },
72                (ip.b1 - ip.b0) as usize,
73                ihd,
74                true,
75            );
76        }
77
78        let any_selective = (base_pos + nb) / ratio > kblk;
79
80        if any_selective {
81            self.dispatch(
82                enc,
83                &self.pipes.index_q,
84                |e| {
85                    self.bind(e, 0, &s.iqk, 0);
86                    self.bind(e, 1, &s.iq, 0);
87                    self.bind(e, 2, &self.dense, a.iqn.0);
88                    set_bytes(e, 3, &ip);
89                },
90                nb * inh,
91                ihd,
92                true,
93            );
94            self.dispatch(
95                enc,
96                &self.pipes.index_score,
97                |e| {
98                    self.bind(e, 0, &s.iq, 0);
99                    self.bind(e, 1, &a.blk, 0);
100                    self.bind(e, 2, &s.bscore, 0);
101                    set_bytes(e, 3, &ip);
102                },
103                nb * max_blocks.div_ceil(256),
104                256,
105                true,
106            );
107            self.dispatch(
108                enc,
109                &self.pipes.index_select,
110                |e| {
111                    self.bind(e, 0, &s.bscore, 0);
112                    self.bind(e, 1, &s.vis, 0);
113                    self.bind(e, 2, &s.nvis, 0);
114                    set_bytes(e, 3, &ip);
115                    self.bind(e, 4, &s.vmask, 0);
116                },
117                nb,
118                1024,
119                true,
120            );
121        }
122
123        let qp = QkRopeParams {
124            n_heads: c.num_attention_heads as u32,
125            head_dim: hd,
126            stride: 2 * hd,
127            rot,
128            pos: base_pos as u32,
129            theta: c.rope_parameters.rope_theta as f32,
130            eps: c.rms_norm_eps as f32,
131        };
132
133        self.dispatch(
134            enc,
135            &self.pipes.qk_norm_rope_b,
136            |e| {
137                self.bind(e, 0, &s.qg, 0);
138                self.bind(e, 1, &self.dense, a.qn.0);
139                set_bytes(e, 2, &qp);
140                set_bytes(e, 3, &nbu);
141            },
142            (nb * c.num_attention_heads).div_ceil(4),
143            128,
144            true,
145        );
146
147        let kp = QkRopeParams {
148            n_heads: c.num_key_value_heads as u32,
149            stride: hd,
150            ..qp
151        };
152
153        self.dispatch(
154            enc,
155            &self.pipes.qk_norm_rope_b,
156            |e| {
157                self.bind(e, 0, &s.k, 0);
158                self.bind(e, 1, &self.dense, a.kn.0);
159                set_bytes(e, 2, &kp);
160                set_bytes(e, 3, &nbu);
161            },
162            (nb * c.num_key_value_heads).div_ceil(4),
163            128,
164            true,
165        );
166
167        let kq = KvQParams {
168            row: kv_row as u32,
169            t0: base_pos as u32,
170            nb: nbu,
171            scale_off: kv_q8_side(self.max_t, kv_row).0 as u32,
172        };
173        let sgs = nb * (kv_row / 32) * 2;
174
175        self.dispatch(
176            enc,
177            &self.pipes.kv_append_q8,
178            |e| {
179                self.bind(e, 0, &s.k, 0);
180                self.bind(e, 1, &s.v, 0);
181                self.bind(e, 2, &a.kc, 0);
182                self.bind(e, 3, &a.vc, 0);
183                set_bytes(e, 4, &kq);
184            },
185            sgs.div_ceil(4),
186            128,
187            true,
188        );
189
190        // Split-T flash decode per row over its causal prefix, or over the
191        // indexer's visible token list once the row is past the budget.
192        let n_rep = c.num_attention_heads / c.num_key_value_heads;
193
194        for b in 0..nb {
195            let t_len = base_pos + b + 1;
196            let blocks = t_len / ratio;
197            let selective = blocks > kblk;
198            let n_vis = if selective {
199                kblk * ratio + (t_len - blocks * ratio)
200            } else {
201                t_len
202            };
203            let n_chunks = n_vis.div_ceil(32);
204            let n_wg = n_chunks.div_ceil(3).clamp(1, ATTN_MAX_WG);
205            let ap = AttnPartParams {
206                n_heads: c.num_attention_heads as u32,
207                n_kv: c.num_key_value_heads as u32,
208                head_dim: hd,
209                t_len: n_vis as u32,
210                q_stride: 2 * hd,
211                q_off: (b * c.num_attention_heads * 2 * c.head_dim) as u32,
212                max_blk: self.max_t.div_ceil(ATTN_TB).max(ATTN_MAX_WG) as u32,
213                scale: (c.head_dim as f32).powf(-0.5) * std::f32::consts::LOG2_E,
214            };
215            let n_wg_u = n_wg as u32;
216            let scale_off = kq.scale_off;
217            let pipe = if selective {
218                &self.pipes.attn_sel
219            } else {
220                &self.pipes.attn_part2_q8
221            };
222
223            self.dispatch(
224                enc,
225                pipe,
226                |e| {
227                    self.bind(e, 0, &s.qg, 0);
228                    self.bind(e, 1, &a.kc, 0);
229                    self.bind(e, 2, &a.vc, 0);
230                    self.bind(e, 3, &s.attn_parts, 0);
231                    set_bytes(e, 4, &ap);
232                    set_bytes(e, 5, &n_wg_u);
233                    set_bytes(e, 6, &scale_off);
234
235                    if selective {
236                        self.bind(e, 7, &s.vis, b * vis_stride * 4);
237                    }
238                },
239                c.num_key_value_heads * n_wg,
240                32 * n_rep,
241                true,
242            );
243
244            let nblk = n_wg as u32;
245            let out_off = (b * c.num_attention_heads * c.head_dim) as u32;
246
247            self.dispatch(
248                enc,
249                &self.pipes.attn_combine,
250                |e| {
251                    self.bind(e, 0, &s.qg, 0);
252                    self.bind(e, 1, &s.attn_parts, 0);
253                    self.bind(e, 2, &s.attn_out, 0);
254                    set_bytes(e, 3, &ap);
255                    set_bytes(e, 4, &nblk);
256                    set_bytes(e, 5, &out_off);
257                },
258                c.num_attention_heads,
259                256,
260                true,
261            );
262        }
263
264        self.prep_h(
265            enc,
266            &s.attn_out,
267            0,
268            (c.num_attention_heads * c.head_dim) as u32,
269            nb,
270            &s.hc.h1,
271        );
272        self.qmv_h(enc, &a.o, &s.mix_out, nb, &s.hc.h1);
273    }
274}