Skip to main content

teeny_kernels/nn/attention/
flash_attn2.rs

1/*
2 * Copyright (c) 2026 Teenygrad.
3 *
4 * Licensed under the Apache License, Version 2.0 (the "License");
5 * you may not use this file except in compliance with the License.
6 * You may obtain a copy of the License at
7 *
8 *   http://www.apache.org/licenses/LICENSE-2.0
9 *
10 * Unless required by applicable law or agreed to in writing, software
11 * distributed under the License is distributed on an "AS IS" BASIS,
12 * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13 * See the License for the specific language governing permissions and
14 * limitations under the License.
15 */
16
17//! Flash Attention 2 — forward and backward kernels.
18//!
19//! **Layout**: all tensors are stored as `[BATCH * N_HEADS, N_CTX, HEAD_DIM]`
20//! row-major (contiguous).  The caller is responsible for reshaping
21//! `[B, H, N, D]` PyTorch tensors to this flat 3-D layout before calling.
22//!
23//! **Algorithm**: each CTA processes one `(batch, head, q_row)` triple.
24//! The kernel iterates over all `N_CTX_K` key/value rows with the online
25//! softmax recurrence (Flash Attention paper, Dao et al. 2022/2023), so
26//! the full `N_CTX_Q × N_CTX_K` attention matrix is never materialised.
27//! Memory is O(N_CTX × HEAD_DIM) per CTA rather than O(N_CTX²).
28//!
29//! HEAD_DIM must be a power of two and is a compile-time const so the
30//! inner vector loads are always fully unmasked.
31
32use 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// ── Forward ───────────────────────────────────────────────────────────────────
41
42/// Flash Attention 2 forward pass.
43///
44/// Inputs  (all `[BH, N_CTX, HEAD_DIM]` row-major, where `BH = BATCH * N_HEADS`):
45///   `q_ptr`, `k_ptr`, `v_ptr`
46///
47/// Outputs:
48///   `o_ptr`  — attention output   `[BH, N_CTX_Q, HEAD_DIM]`
49///   `l_ptr`  — log-sum-exp        `[BH, N_CTX_Q]`  (saved for backward)
50///
51/// Grid: `(N_CTX_Q, BH, 1)` — one CTA per `(batch_head, q_row)` pair.
52#[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, // 1 / sqrt(HEAD_DIM)
62    neg_inf: f32,       // f32::NEG_INFINITY — passed explicitly (no_core has no float constants)
63) 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); // query-row index  [0, n_ctx_q)
69    let pid_bh = T::program_id(Axis::Y); // (batch, head)    [0, BH)
70
71    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    // HEAD_DIM lane offsets — no masking needed (HEAD_DIM is a power of two).
77    let d = T::arange(0, HEAD_DIM);
78
79    // Load Q[pid_bh, pid_m, :]
80    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    // Online-softmax running state — all kept as [HEAD_DIM] tensors so that
92    // all scf.for iter-args have the same shape (Triton requires this).
93    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        // Scaled dot-product: score = sum(q · k) * scale  → [HD] (scalar replicated)
123        let qk = T::sum(q_vec * k_vec, Some(0), true) * scale_t;
124
125        // Online softmax recurrence
126        let m_new = T::maximum(m_i, qk); // [HD] running max (all elements equal)
127        let exp_diff = T::exp(m_i - m_new); // [HD] correction factor
128        let p = T::exp(qk - m_new); // [HD] unnorm weight for this k
129
130        l_i = exp_diff * l_i + p;
131        acc = exp_diff * acc + p * v_vec;
132        m_i = m_new;
133    }
134
135    // Normalise output and compute logsumexp for backward.
136    let o_row = acc / l_i; // [HD] / [HD] → [HD]
137    // All elements of m_i and l_i are equal (replicated scalar). Sum and divide
138    // by HEAD_DIM to recover the scalar as tensor<1xD> for the l_ptr store.
139    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// ── Backward — dQ ─────────────────────────────────────────────────────────────
161
162/// Flash Attention 2 backward: computes `dQ`.
163///
164/// For each query row `q`, iterates over all key rows `k` and accumulates:
165/// ```text
166/// dQ_q += dS_{qk} * K_k * scale
167/// where dS_{qk} = p_{qk} * (dO_q · V_k − D_q)
168///       p_{qk}  = exp(Q_q · K_k * scale − L_q)   (recomputed attention)
169///       D_q     = sum(O_q * dO_q)                 (per-row scalar)
170/// ```
171///
172/// Grid: `(N_CTX_Q, BH, 1)` — same grid shape as the forward pass.
173#[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    // D_q = rowsum(O_q * dO_q)  — scalar
232    let d_q = T::sum(o_vec * do_vec, Some(0), false);
233
234    // Load logsumexp L_q (scalar stored as 1-element vec); reduce to scalar.
235    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        // Recompute attention weight: p = exp(qk * scale - L_q) — [HEAD_DIM] (scalar replicated)
274        let qk = T::sum(q_vec * k_vec, Some(0), false) * scale_t;
275        let p = T::exp(qk - l_q);
276
277        // dS = p * (dO · V_k - D_q)
278        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 += dS * K_k * scale
282        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// ── Backward — dK / dV ────────────────────────────────────────────────────────
296
297/// Flash Attention 2 backward: computes `dK` and `dV` for one key row.
298///
299/// For each key row `n`, iterates over all query rows `m` and accumulates:
300/// ```text
301/// dV_n += p_{mn} * dO_m
302/// dK_n += dS_{mn} * Q_m * scale
303/// where p_{mn}  = exp(Q_m · K_n * scale − L_m)
304///       dS_{mn} = p_{mn} * (dO_m · V_n − D_m)
305///       D_m     = sum(O_m * dO_m)
306/// ```
307///
308/// Each CTA owns an exclusive key row so no atomic operations are needed.
309///
310/// Grid: `(N_CTX_K, BH, 1)`.
311#[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); // key-row index  [0, n_ctx_k)
330    let pid_bh = T::program_id(Axis::Y); // (batch, head)  [0, BH)
331
332    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    // Load K_n and V_n — fixed for this CTA.
340    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        // Load logsumexp L_m; reduce tensor<1xD> to scalar.
399        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        // D_m = rowsum(O_m * dO_m) — scalar
412        let d_m = T::sum(o_vec_m * do_vec_m, Some(0), false);
413
414        // Recompute p_{mn} = exp(Q_m · K_n * scale - L_m) — [HEAD_DIM] (scalar replicated)
415        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 += p * dO_m
419        dv_acc = dv_acc + p * do_vec_m;
420
421        // dS = p * (dO_m · V_n - D_m)
422        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 += dS * Q_m * scale
426        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
447// ── Op wrapper ────────────────────────────────────────────────────────────────
448
449pub 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}