Skip to main content

vision_rs/models/yolo/kernels/attention/
flash_attn2.rs

1/*
2 * Copyright 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
18//! Flash Attention 2 — forward and backward kernels.
19//!
20//! **Layout**: all tensors are stored as `[BATCH * N_HEADS, N_CTX, HEAD_DIM]`
21//! row-major (contiguous).  The caller is responsible for reshaping
22//! `[B, H, N, D]` PyTorch tensors to this flat 3-D layout before calling.
23//!
24//! **Algorithm**: each CTA processes one `(batch, head, q_row)` triple.
25//! The kernel iterates over all `N_CTX_K` key/value rows with the online
26//! softmax recurrence (Flash Attention paper, Dao et al. 2022/2023), so
27//! the full `N_CTX_Q × N_CTX_K` attention matrix is never materialised.
28//! Memory is O(N_CTX × HEAD_DIM) per CTA rather than O(N_CTX²).
29//!
30//! HEAD_DIM must be a power of two and is a compile-time const so the
31//! inner vector loads are always fully unmasked.
32
33use 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// ── Forward ───────────────────────────────────────────────────────────────────
42
43/// Flash Attention 2 forward pass.
44///
45/// Inputs  (all `[BH, N_CTX, HEAD_DIM]` row-major, where `BH = BATCH * N_HEADS`):
46///   `q_ptr`, `k_ptr`, `v_ptr`
47///
48/// Outputs:
49///   `o_ptr`  — attention output   `[BH, N_CTX_Q, HEAD_DIM]`
50///   `l_ptr`  — log-sum-exp        `[BH, N_CTX_Q]`  (saved for backward)
51///
52/// Grid: `(N_CTX_Q, BH, 1)` — one CTA per `(batch_head, q_row)` pair.
53#[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, // 1 / sqrt(HEAD_DIM)
63    neg_inf: f32,       // f32::NEG_INFINITY — passed explicitly (no_core has no float constants)
64) 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); // query-row index  [0, n_ctx_q)
70    let pid_bh = T::program_id(Axis::Y); // (batch, head)    [0, BH)
71
72    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    // HEAD_DIM lane offsets — no masking needed (HEAD_DIM is a power of two).
78    let d = T::arange(0, HEAD_DIM);
79
80    // Load Q[pid_bh, pid_m, :]
81    let q_vec = T::load(
82        q_ptr.add_offsets(d + q_row_base),
83        None, None, &[], None, None, None, false,
84    );
85
86    // Online-softmax running state — all kept as [HEAD_DIM] tensors so that
87    // all scf.for iter-args have the same shape (Triton requires this).
88    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        // Scaled dot-product: score = sum(q · k) * scale  → [HD] (scalar replicated)
100        let qk = T::sum(q_vec * k_vec, Some(0), true) * scale_t;
101
102        // Online softmax recurrence
103        let m_new = T::maximum(m_i, qk); // [HD] running max (all elements equal)
104        let exp_diff = T::exp(m_i - m_new); // [HD] correction factor
105        let p = T::exp(qk - m_new); // [HD] unnorm weight for this k
106
107        l_i = exp_diff * l_i + p;
108        acc = exp_diff * acc + p * v_vec;
109        m_i = m_new;
110    }
111
112    // Normalise output and compute logsumexp for backward.
113    let o_row = acc / l_i; // [HD] / [HD] → [HD]
114    // All elements of m_i and l_i are equal (replicated scalar). Sum and divide
115    // by HEAD_DIM to recover the scalar as tensor<1xf32> for the l_ptr store.
116    let l_save_sum = T::sum(m_i + T::log(l_i), Some(0), false); // scalar f32
117    let l_save = l_save_sum / T::full(&[1], D::from_f64(HEAD_DIM as f64)); // tensor<1xD>
118
119    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// ── Backward — dQ ─────────────────────────────────────────────────────────────
124
125/// Flash Attention 2 backward: computes `dQ`.
126///
127/// Grid: `(N_CTX_Q, BH, 1)` — same grid shape as the forward pass.
128#[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    // D_q = rowsum(O_q * dO_q)  — scalar
160    let d_q = T::sum(o_vec * do_vec, Some(0), false);
161
162    // Load logsumexp L_q (scalar stored as 1-element vec); reduce to scalar.
163    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); // scalar
165
166    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        // Recompute attention weight: p = exp(qk * scale - L_q)
175        let qk = T::sum(q_vec * k_vec, Some(0), false) * scale_t;
176        let p = T::exp(qk - l_q);
177
178        // dS = p * (dO · V_k - D_q)
179        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 += dS * K_k * scale
183        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// ── Backward — dK / dV ────────────────────────────────────────────────────────
190
191/// Flash Attention 2 backward: computes `dK` and `dV` for one key row.
192///
193/// Grid: `(N_CTX_K, BH, 1)`.
194#[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); // key-row index  [0, n_ctx_k)
213    let pid_bh = T::program_id(Axis::Y); // (batch, head)  [0, BH)
214
215    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    // Load K_n and V_n — fixed for this CTA.
223    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        // Load logsumexp L_m; reduce tensor<1xf32> to scalar.
237        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); // scalar f32
239
240        // D_m = rowsum(O_m * dO_m) — scalar
241        let d_m = T::sum(o_vec_m * do_vec_m, Some(0), false);
242
243        // Recompute p_{mn} = exp(Q_m · K_n * scale - L_m)
244        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 += p * dO_m
248        dv_acc = dv_acc + p * do_vec_m;
249
250        // dS = p * (dO_m · V_n - D_m)
251        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 += dS * Q_m * scale
255        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
262// ── Op wrapper ────────────────────────────────────────────────────────────────
263
264/// Bundles the forward and backward Flash Attention 2 kernels for a given head dimension.
265pub struct FlashAttention2Op<'a, D: Float + Send + Sync + 'static> {
266    /// The forward-pass kernel.
267    pub forward: FlashAttention2Forward<D>,
268    /// The backward-pass kernel computing `dQ`.
269    pub backward_dq: FlashAttention2BackwardDq<D>,
270    /// The backward-pass kernel computing `dK` and `dV`.
271    pub backward_dkv: FlashAttention2BackwardDkv<D>,
272    _marker: PhantomData<&'a ()>,
273}
274
275impl<'a, D: Float + Send + Sync + 'static> FlashAttention2Op<'a, D> {
276    /// Constructs the forward/backward kernel set for the given head dimension.
277    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}