1use core::marker::PhantomData;
33use teeny_core::dtype::Float;
34use teeny_macros::kernel;
35use teeny_triton::triton::{
36 types::{AddOffsets, Comparison, Tensor},
37 *,
38};
39
40#[kernel]
53pub fn flash_attention2_forward<T: Triton, D: Float, const HEAD_DIM: i32>(
54 q_ptr: T::Pointer<D>,
55 k_ptr: T::Pointer<D>,
56 v_ptr: T::Pointer<D>,
57 o_ptr: T::Pointer<D>,
58 l_ptr: T::Pointer<D>,
59 n_ctx_q: i32,
60 n_ctx_k: i32,
61 softmax_scale: f32, neg_inf: f32, ) where
64 T::I32Tensor: Tensor<i32, 1>,
65 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
66 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
67{
68 let pid_m = T::program_id(Axis::X); let pid_bh = T::program_id(Axis::Y); let kv_bh_base = pid_bh * n_ctx_k * HEAD_DIM;
72 let q_row_base = pid_bh * n_ctx_q * HEAD_DIM + pid_m * HEAD_DIM;
73 let o_row_base = pid_bh * n_ctx_q * HEAD_DIM + pid_m * HEAD_DIM;
74 let l_row_base = pid_bh * n_ctx_q + pid_m;
75
76 let d = T::arange(0, HEAD_DIM);
78
79 let q_vec = T::load(
81 q_ptr.add_offsets(d + q_row_base),
82 None,
83 None,
84 &[],
85 None,
86 None,
87 None,
88 false,
89 );
90
91 let mut acc = T::zeros::<D>(&[HEAD_DIM]);
94 let mut m_i = T::full(&[HEAD_DIM], D::from_f64(neg_inf as f64));
95 let mut l_i = T::zeros::<D>(&[HEAD_DIM]);
96 let scale_t = T::full(&[HEAD_DIM], D::from_f64(softmax_scale as f64));
97
98 for k_row in 0..n_ctx_k {
99 let kv_row_base = kv_bh_base + k_row * HEAD_DIM;
100
101 let k_vec = T::load(
102 k_ptr.add_offsets(d + kv_row_base),
103 None,
104 None,
105 &[],
106 None,
107 None,
108 None,
109 false,
110 );
111 let v_vec = T::load(
112 v_ptr.add_offsets(d + kv_row_base),
113 None,
114 None,
115 &[],
116 None,
117 None,
118 None,
119 false,
120 );
121
122 let qk = T::sum(q_vec * k_vec, Some(0), true) * scale_t;
124
125 let m_new = T::maximum(m_i, qk); let exp_diff = T::exp(m_i - m_new); let p = T::exp(qk - m_new); l_i = exp_diff * l_i + p;
131 acc = exp_diff * acc + p * v_vec;
132 m_i = m_new;
133 }
134
135 let o_row = acc / l_i; let l_save_sum = T::sum(m_i + T::log(l_i), Some(0), false);
140 let l_save = l_save_sum / T::full(&[1], D::from_f64(HEAD_DIM as f64));
141
142 T::store(
143 o_ptr.add_offsets(d + o_row_base),
144 o_row,
145 None,
146 &[],
147 None,
148 None,
149 );
150 T::store(
151 l_ptr.add_offsets(T::arange(0, 1) + l_row_base),
152 l_save,
153 None,
154 &[],
155 None,
156 None,
157 );
158}
159
160#[kernel]
174pub fn flash_attention2_backward_dq<T: Triton, D: Float, const HEAD_DIM: i32>(
175 q_ptr: T::Pointer<D>,
176 k_ptr: T::Pointer<D>,
177 v_ptr: T::Pointer<D>,
178 o_ptr: T::Pointer<D>,
179 do_ptr: T::Pointer<D>,
180 l_ptr: T::Pointer<D>,
181 dq_ptr: T::Pointer<D>,
182 n_ctx_q: i32,
183 n_ctx_k: i32,
184 softmax_scale: f32,
185) where
186 T::I32Tensor: Tensor<i32, 1>,
187 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
188 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
189{
190 let pid_m = T::program_id(Axis::X);
191 let pid_bh = T::program_id(Axis::Y);
192
193 let q_row_base = pid_bh * n_ctx_q * HEAD_DIM + pid_m * HEAD_DIM;
194 let kv_bh_base = pid_bh * n_ctx_k * HEAD_DIM;
195 let l_row_base = pid_bh * n_ctx_q + pid_m;
196
197 let d = T::arange(0, HEAD_DIM);
198 let scale_t = T::full(&[HEAD_DIM], D::from_f64(softmax_scale as f64));
199
200 let q_vec = T::load(
201 q_ptr.add_offsets(d + q_row_base),
202 None,
203 None,
204 &[],
205 None,
206 None,
207 None,
208 false,
209 );
210 let o_vec = T::load(
211 o_ptr.add_offsets(d + q_row_base),
212 None,
213 None,
214 &[],
215 None,
216 None,
217 None,
218 false,
219 );
220 let do_vec = T::load(
221 do_ptr.add_offsets(d + q_row_base),
222 None,
223 None,
224 &[],
225 None,
226 None,
227 None,
228 false,
229 );
230
231 let d_q = T::sum(o_vec * do_vec, Some(0), false);
233
234 let l_q_raw = T::load(
236 l_ptr.add_offsets(T::arange(0, 1) + l_row_base),
237 None,
238 None,
239 &[],
240 None,
241 None,
242 None,
243 false,
244 );
245 let l_q = T::sum(l_q_raw, Some(0), false);
246
247 let mut dq_acc = T::zeros::<D>(&[HEAD_DIM]);
248
249 for k_row in 0..n_ctx_k {
250 let kv_row_base = kv_bh_base + k_row * HEAD_DIM;
251
252 let k_vec = T::load(
253 k_ptr.add_offsets(d + kv_row_base),
254 None,
255 None,
256 &[],
257 None,
258 None,
259 None,
260 false,
261 );
262 let v_vec = T::load(
263 v_ptr.add_offsets(d + kv_row_base),
264 None,
265 None,
266 &[],
267 None,
268 None,
269 None,
270 false,
271 );
272
273 let qk = T::sum(q_vec * k_vec, Some(0), false) * scale_t;
275 let p = T::exp(qk - l_q);
276
277 let do_dot_v = T::sum(do_vec * v_vec, Some(0), false);
279 let ds = p * (do_dot_v - d_q);
280
281 dq_acc = dq_acc + ds * k_vec * scale_t;
283 }
284
285 T::store(
286 dq_ptr.add_offsets(d + q_row_base),
287 dq_acc,
288 None,
289 &[],
290 None,
291 None,
292 );
293}
294
295#[kernel]
312pub fn flash_attention2_backward_dkv<T: Triton, D: Float, const HEAD_DIM: i32>(
313 q_ptr: T::Pointer<D>,
314 k_ptr: T::Pointer<D>,
315 v_ptr: T::Pointer<D>,
316 o_ptr: T::Pointer<D>,
317 do_ptr: T::Pointer<D>,
318 l_ptr: T::Pointer<D>,
319 dk_ptr: T::Pointer<D>,
320 dv_ptr: T::Pointer<D>,
321 n_ctx_q: i32,
322 n_ctx_k: i32,
323 softmax_scale: f32,
324) where
325 T::I32Tensor: Tensor<i32, 1>,
326 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
327 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
328{
329 let pid_n = T::program_id(Axis::X); let pid_bh = T::program_id(Axis::Y); let q_bh_base = pid_bh * n_ctx_q * HEAD_DIM;
333 let kv_row_base = pid_bh * n_ctx_k * HEAD_DIM + pid_n * HEAD_DIM;
334 let l_bh_base = pid_bh * n_ctx_q;
335
336 let d = T::arange(0, HEAD_DIM);
337 let scale_t = T::full(&[HEAD_DIM], D::from_f64(softmax_scale as f64));
338
339 let k_vec = T::load(
341 k_ptr.add_offsets(d + kv_row_base),
342 None,
343 None,
344 &[],
345 None,
346 None,
347 None,
348 false,
349 );
350 let v_vec = T::load(
351 v_ptr.add_offsets(d + kv_row_base),
352 None,
353 None,
354 &[],
355 None,
356 None,
357 None,
358 false,
359 );
360
361 let mut dk_acc = T::zeros::<D>(&[HEAD_DIM]);
362 let mut dv_acc = T::zeros::<D>(&[HEAD_DIM]);
363
364 for q_row in 0..n_ctx_q {
365 let q_row_base = q_bh_base + q_row * HEAD_DIM;
366 let l_row_base = l_bh_base + q_row;
367
368 let q_vec_m = T::load(
369 q_ptr.add_offsets(d + q_row_base),
370 None,
371 None,
372 &[],
373 None,
374 None,
375 None,
376 false,
377 );
378 let o_vec_m = T::load(
379 o_ptr.add_offsets(d + q_row_base),
380 None,
381 None,
382 &[],
383 None,
384 None,
385 None,
386 false,
387 );
388 let do_vec_m = T::load(
389 do_ptr.add_offsets(d + q_row_base),
390 None,
391 None,
392 &[],
393 None,
394 None,
395 None,
396 false,
397 );
398 let l_m_raw = T::load(
400 l_ptr.add_offsets(T::arange(0, 1) + l_row_base),
401 None,
402 None,
403 &[],
404 None,
405 None,
406 None,
407 false,
408 );
409 let l_m = T::sum(l_m_raw, Some(0), false);
410
411 let d_m = T::sum(o_vec_m * do_vec_m, Some(0), false);
413
414 let qk = T::sum(q_vec_m * k_vec, Some(0), false) * scale_t;
416 let p = T::exp(qk - l_m);
417
418 dv_acc = dv_acc + p * do_vec_m;
420
421 let do_dot_v = T::sum(do_vec_m * v_vec, Some(0), false);
423 let ds = p * (do_dot_v - d_m);
424
425 dk_acc = dk_acc + ds * q_vec_m * scale_t;
427 }
428
429 T::store(
430 dk_ptr.add_offsets(d + kv_row_base),
431 dk_acc,
432 None,
433 &[],
434 None,
435 None,
436 );
437 T::store(
438 dv_ptr.add_offsets(d + kv_row_base),
439 dv_acc,
440 None,
441 &[],
442 None,
443 None,
444 );
445}
446
447pub struct FlashAttention2Op<'a, D: Float + Send + Sync + 'static> {
450 pub forward: FlashAttention2Forward<D>,
451 pub backward_dq: FlashAttention2BackwardDq<D>,
452 pub backward_dkv: FlashAttention2BackwardDkv<D>,
453 _marker: PhantomData<&'a ()>,
454}
455
456impl<'a, D: Float + Send + Sync + 'static> FlashAttention2Op<'a, D> {
457 pub fn new(head_dim: i32) -> Self {
458 Self {
459 forward: FlashAttention2Forward::<D>::new(head_dim),
460 backward_dq: FlashAttention2BackwardDq::<D>::new(head_dim),
461 backward_dkv: FlashAttention2BackwardDkv::<D>::new(head_dim),
462 _marker: PhantomData,
463 }
464 }
465}