1use core::marker::PhantomData;
34use teeny_core::dtype::Float;
35use teeny_macros::kernel;
36use teeny_triton::triton::{
37 types::{AddOffsets, Comparison, Tensor},
38 *,
39};
40
41#[kernel]
54pub fn flash_attention2_forward<T: Triton, D: Float, const HEAD_DIM: i32>(
55 q_ptr: T::Pointer<D>,
56 k_ptr: T::Pointer<D>,
57 v_ptr: T::Pointer<D>,
58 o_ptr: T::Pointer<D>,
59 l_ptr: T::Pointer<D>,
60 n_ctx_q: i32,
61 n_ctx_k: i32,
62 softmax_scale: f32, neg_inf: f32, ) where
65 T::I32Tensor: Tensor<i32, 1>,
66 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
67 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
68{
69 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;
73 let q_row_base = pid_bh * n_ctx_q * HEAD_DIM + pid_m * HEAD_DIM;
74 let o_row_base = pid_bh * n_ctx_q * HEAD_DIM + pid_m * HEAD_DIM;
75 let l_row_base = pid_bh * n_ctx_q + pid_m;
76
77 let d = T::arange(0, HEAD_DIM);
79
80 let q_vec = T::load(
82 q_ptr.add_offsets(d + q_row_base),
83 None, None, &[], None, None, None, false,
84 );
85
86 let mut acc = T::zeros::<D>(&[HEAD_DIM]);
89 let mut m_i = T::full(&[HEAD_DIM], D::from_f64(neg_inf as f64));
90 let mut l_i = T::zeros::<D>(&[HEAD_DIM]);
91 let scale_t = T::full(&[HEAD_DIM], D::from_f64(softmax_scale as f64));
92
93 for k_row in 0..n_ctx_k {
94 let kv_row_base = kv_bh_base + k_row * HEAD_DIM;
95
96 let k_vec = T::load(k_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
97 let v_vec = T::load(v_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
98
99 let qk = T::sum(q_vec * k_vec, Some(0), true) * scale_t;
101
102 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;
108 acc = exp_diff * acc + p * v_vec;
109 m_i = m_new;
110 }
111
112 let o_row = acc / l_i; let l_save_sum = T::sum(m_i + T::log(l_i), Some(0), false); let l_save = l_save_sum / T::full(&[1], D::from_f64(HEAD_DIM as f64)); T::store(o_ptr.add_offsets(d + o_row_base), o_row, None, &[], None, None);
120 T::store(l_ptr.add_offsets(T::arange(0, 1) + l_row_base), l_save, None, &[], None, None);
121}
122
123#[kernel]
129pub fn flash_attention2_backward_dq<T: Triton, D: Float, const HEAD_DIM: i32>(
130 q_ptr: T::Pointer<D>,
131 k_ptr: T::Pointer<D>,
132 v_ptr: T::Pointer<D>,
133 o_ptr: T::Pointer<D>,
134 do_ptr: T::Pointer<D>,
135 l_ptr: T::Pointer<D>,
136 dq_ptr: T::Pointer<D>,
137 n_ctx_q: i32,
138 n_ctx_k: i32,
139 softmax_scale: f32,
140) where
141 T::I32Tensor: Tensor<i32, 1>,
142 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
143 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
144{
145 let pid_m = T::program_id(Axis::X);
146 let pid_bh = T::program_id(Axis::Y);
147
148 let q_row_base = pid_bh * n_ctx_q * HEAD_DIM + pid_m * HEAD_DIM;
149 let kv_bh_base = pid_bh * n_ctx_k * HEAD_DIM;
150 let l_row_base = pid_bh * n_ctx_q + pid_m;
151
152 let d = T::arange(0, HEAD_DIM);
153 let scale_t = T::full(&[HEAD_DIM], D::from_f64(softmax_scale as f64));
154
155 let q_vec = T::load(q_ptr.add_offsets(d + q_row_base), None, None, &[], None, None, None, false);
156 let o_vec = T::load(o_ptr.add_offsets(d + q_row_base), None, None, &[], None, None, None, false);
157 let do_vec = T::load(do_ptr.add_offsets(d + q_row_base), None, None, &[], None, None, None, false);
158
159 let d_q = T::sum(o_vec * do_vec, Some(0), false);
161
162 let l_q_raw = T::load(l_ptr.add_offsets(T::arange(0, 1) + l_row_base), None, None, &[], None, None, None, false);
164 let l_q = T::sum(l_q_raw, Some(0), false); let mut dq_acc = T::zeros::<D>(&[HEAD_DIM]);
167
168 for k_row in 0..n_ctx_k {
169 let kv_row_base = kv_bh_base + k_row * HEAD_DIM;
170
171 let k_vec = T::load(k_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
172 let v_vec = T::load(v_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
173
174 let qk = T::sum(q_vec * k_vec, Some(0), false) * scale_t;
176 let p = T::exp(qk - l_q);
177
178 let do_dot_v = T::sum(do_vec * v_vec, Some(0), false);
180 let ds = p * (do_dot_v - d_q);
181
182 dq_acc = dq_acc + ds * k_vec * scale_t;
184 }
185
186 T::store(dq_ptr.add_offsets(d + q_row_base), dq_acc, None, &[], None, None);
187}
188
189#[kernel]
195pub fn flash_attention2_backward_dkv<T: Triton, D: Float, const HEAD_DIM: i32>(
196 q_ptr: T::Pointer<D>,
197 k_ptr: T::Pointer<D>,
198 v_ptr: T::Pointer<D>,
199 o_ptr: T::Pointer<D>,
200 do_ptr: T::Pointer<D>,
201 l_ptr: T::Pointer<D>,
202 dk_ptr: T::Pointer<D>,
203 dv_ptr: T::Pointer<D>,
204 n_ctx_q: i32,
205 n_ctx_k: i32,
206 softmax_scale: f32,
207) where
208 T::I32Tensor: Tensor<i32, 1>,
209 T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
210 T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
211{
212 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;
216 let kv_row_base = pid_bh * n_ctx_k * HEAD_DIM + pid_n * HEAD_DIM;
217 let l_bh_base = pid_bh * n_ctx_q;
218
219 let d = T::arange(0, HEAD_DIM);
220 let scale_t = T::full(&[HEAD_DIM], D::from_f64(softmax_scale as f64));
221
222 let k_vec = T::load(k_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
224 let v_vec = T::load(v_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
225
226 let mut dk_acc = T::zeros::<D>(&[HEAD_DIM]);
227 let mut dv_acc = T::zeros::<D>(&[HEAD_DIM]);
228
229 for q_row in 0..n_ctx_q {
230 let q_row_base = q_bh_base + q_row * HEAD_DIM;
231 let l_row_base = l_bh_base + q_row;
232
233 let q_vec_m = T::load(q_ptr.add_offsets(d + q_row_base), None, None, &[], None, None, None, false);
234 let o_vec_m = T::load(o_ptr.add_offsets(d + q_row_base), None, None, &[], None, None, None, false);
235 let do_vec_m = T::load(do_ptr.add_offsets(d + q_row_base), None, None, &[], None, None, None, false);
236 let l_m_raw = T::load(l_ptr.add_offsets(T::arange(0, 1) + l_row_base), None, None, &[], None, None, None, false);
238 let l_m = T::sum(l_m_raw, Some(0), false); let d_m = T::sum(o_vec_m * do_vec_m, Some(0), false);
242
243 let qk = T::sum(q_vec_m * k_vec, Some(0), false) * scale_t;
245 let p = T::exp(qk - l_m);
246
247 dv_acc = dv_acc + p * do_vec_m;
249
250 let do_dot_v = T::sum(do_vec_m * v_vec, Some(0), false);
252 let ds = p * (do_dot_v - d_m);
253
254 dk_acc = dk_acc + ds * q_vec_m * scale_t;
256 }
257
258 T::store(dk_ptr.add_offsets(d + kv_row_base), dk_acc, None, &[], None, None);
259 T::store(dv_ptr.add_offsets(d + kv_row_base), dv_acc, None, &[], None, None);
260}
261
262pub struct FlashAttention2Op<'a, D: Float + Send + Sync + 'static> {
266 pub forward: FlashAttention2Forward<D>,
268 pub backward_dq: FlashAttention2BackwardDq<D>,
270 pub backward_dkv: FlashAttention2BackwardDkv<D>,
272 _marker: PhantomData<&'a ()>,
273}
274
275impl<'a, D: Float + Send + Sync + 'static> FlashAttention2Op<'a, D> {
276 pub fn new(head_dim: i32) -> Self {
278 Self {
279 forward: FlashAttention2Forward::<D>::new(head_dim),
280 backward_dq: FlashAttention2BackwardDq::<D>::new(head_dim),
281 backward_dkv: FlashAttention2BackwardDkv::<D>::new(head_dim),
282 _marker: PhantomData,
283 }
284 }
285}