1use super::*;
4
5impl Gpu<'_> {
6 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 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 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}