Skip to main content

vision_rs/models/yolo/kernels/attention/
psa.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//! PSA attention helper kernels and RuntimeOp wrappers.
19//!
20//! Implements the data-rearrangement passes surrounding Flash Attention 2 for
21//! the PSABlock used in YOLO26 C2PSA layers.
22//!
23//! Assumptions (derived from the ultralytics PSABlock / Attention module):
24//!   - `head_dim = c / num_heads = 2 * key_dim`
25//!   - QKV conv output has `qkv_h = num_heads * 4 * KEY_DIM` channels
26//!     stored as NCHW `[B, qkv_h, H, W]`.
27//!   - Per head `h`, channels `[h*4*KEY_DIM : +KEY_DIM]` are Q,
28//!     `[+KEY_DIM : +2*KEY_DIM]` are K, `[+2*KEY_DIM : +3*KEY_DIM]` are V_lo,
29//!     `[+3*KEY_DIM : +4*KEY_DIM]` are V_hi.
30//!
31//! The V split trick lets us run FA2 with `HEAD_DIM = key_dim` twice (once for
32//! V_lo, once for V_hi) instead of needing `HEAD_DIM = head_dim = 2*key_dim`.
33
34#![allow(non_snake_case)]
35
36use core::ffi::c_void;
37
38use teeny_core::dtype::Float;
39use teeny_macros::kernel;
40use teeny_triton::triton::{
41    types::{AddOffsets, Comparison, Tensor},
42    *,
43};
44
45use super::flash_attn2::FlashAttention2Forward;
46
47// ── psa_pack_qkv ─────────────────────────────────────────────────────────────
48
49/// Rearranges the NCHW QKV tensor into packed FA2 format.
50///
51/// Input:  `qkv_ptr`  — `[B, qkv_h, H, W]` NCHW, `qkv_h = num_heads * 4 * KEY_DIM`
52/// Output: `out_ptr`  — flat `[4, BH, N, KEY_DIM]` buffer
53///   - Section 0: Q, Section 1: K, Section 2: V_lo, Section 3: V_hi
54///
55/// Grid: `[4 * BH * N, 1, 1]` — one CTA per (section, bh, n) triple.
56/// Block: `[KEY_DIM, 1, 1]`
57#[kernel]
58pub fn psa_pack_qkv<T: Triton, D: Float, const KEY_DIM: i32>(
59    qkv_ptr: T::Pointer<D>,
60    out_ptr: T::Pointer<D>,
61    qkv_h: i32,    // num_heads * 4 * KEY_DIM
62    H: i32,
63    W: i32,
64    B: i32,
65    num_heads: i32,
66) where
67    T::I32Tensor: Tensor<i32, 1>,
68    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
69    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
70{
71    let pid = T::program_id(Axis::X); // [0, 4 * BH * N)
72    let BH: i32 = B * num_heads;
73    let N: i32 = H * W;
74
75    let section: i32 = pid / (BH * N);
76    let bh: i32 = (pid / N) % BH;
77    let n: i32 = pid % N;
78    let b: i32 = bh / num_heads;
79    let h: i32 = bh % num_heads;
80
81    let d = T::arange(0, KEY_DIM);
82
83    // NCHW source: offset = (h*4*KEY_DIM + section*KEY_DIM + d) * H*W + b*qkv_h*H*W + n
84    let chan_base: i32 = h * 4 * KEY_DIM + section * KEY_DIM;
85    let src_off = (d + chan_base) * (H * W) + (b * qkv_h * H * W + n);
86
87    let x = T::load(qkv_ptr.add_offsets(src_off), None, None, &[], None, None, None, false);
88
89    // Output flat [4, BH, N, KEY_DIM]
90    let dst_base: i32 = section * BH * N * KEY_DIM + bh * N * KEY_DIM + n * KEY_DIM;
91    let dst_off = d + dst_base;
92
93    T::store(out_ptr.add_offsets(dst_off), x, None, &[], None, None);
94}
95
96// ── psa_extract_v_nchw ────────────────────────────────────────────────────────
97
98/// Extracts V channels from QKV NCHW into V NCHW.
99///
100/// Input:  `qkv_ptr`  — `[B, qkv_h, H, W]` NCHW
101/// Output: `v_ptr`    — `[B, c, H, W]` NCHW  (c = num_heads * 2 * KEY_DIM)
102///
103/// Grid: `[BH * N, 1, 1]`.  Block: `[KEY_DIM, 1, 1]`.
104#[kernel]
105pub fn psa_extract_v_nchw<T: Triton, D: Float, const KEY_DIM: i32>(
106    qkv_ptr: T::Pointer<D>,
107    v_ptr: T::Pointer<D>,
108    qkv_h: i32,
109    c: i32,        // num_heads * 2 * KEY_DIM
110    H: i32,
111    W: i32,
112    num_heads: i32,
113) where
114    T::I32Tensor: Tensor<i32, 1>,
115    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
116    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
117{
118    let pid = T::program_id(Axis::X); // [0, BH * N)
119    let N: i32 = H * W;
120    let bh: i32 = pid / N;
121    let n: i32 = pid % N;
122    let b: i32 = bh / num_heads;
123    let h: i32 = bh % num_heads;
124
125    let d = T::arange(0, KEY_DIM);
126
127    let src_lo_base: i32 = h * 4 * KEY_DIM + 2 * KEY_DIM;
128    let src_hi_base: i32 = h * 4 * KEY_DIM + 3 * KEY_DIM;
129    let src_off_lo = (d + src_lo_base) * (H * W) + (b * qkv_h * H * W + n);
130    let src_off_hi = (d + src_hi_base) * (H * W) + (b * qkv_h * H * W + n);
131
132    let x_lo = T::load(qkv_ptr.add_offsets(src_off_lo), None, None, &[], None, None, None, false);
133    let x_hi = T::load(qkv_ptr.add_offsets(src_off_hi), None, None, &[], None, None, None, false);
134
135    let dst_lo_base: i32 = h * 2 * KEY_DIM;
136    let dst_hi_base: i32 = h * 2 * KEY_DIM + KEY_DIM;
137    let dst_off_lo = (d + dst_lo_base) * (H * W) + (b * c * H * W + n);
138    let dst_off_hi = (d + dst_hi_base) * (H * W) + (b * c * H * W + n);
139
140    T::store(v_ptr.add_offsets(dst_off_lo), x_lo, None, &[], None, None);
141    T::store(v_ptr.add_offsets(dst_off_hi), x_hi, None, &[], None, None);
142}
143
144// ── psa_merge_attn_nchw ───────────────────────────────────────────────────────
145
146/// Merges two FA2 outputs (V_lo and V_hi attention results) into NCHW.
147///
148/// Inputs: `lo_ptr`, `hi_ptr` — each `[BH * N * KEY_DIM]` flat
149/// Output: `out_ptr` — `[B, c, H, W]` NCHW  (c = num_heads * 2 * KEY_DIM)
150///
151/// Grid: `[BH * N, 1, 1]`.  Block: `[KEY_DIM, 1, 1]`.
152#[kernel]
153pub fn psa_merge_attn_nchw<T: Triton, D: Float, const KEY_DIM: i32>(
154    lo_ptr: T::Pointer<D>,
155    hi_ptr: T::Pointer<D>,
156    out_ptr: T::Pointer<D>,
157    c: i32,
158    H: i32,
159    W: i32,
160    num_heads: i32,
161) where
162    T::I32Tensor: Tensor<i32, 1>,
163    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
164    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
165{
166    let pid = T::program_id(Axis::X); // [0, BH * N)
167    let N: i32 = H * W;
168    let bh: i32 = pid / N;
169    let n: i32 = pid % N;
170    let b: i32 = bh / num_heads;
171    let h: i32 = bh % num_heads;
172
173    let d = T::arange(0, KEY_DIM);
174
175    // Source: flat [BH, N, KEY_DIM]
176    let src_base: i32 = bh * N * KEY_DIM + n * KEY_DIM;
177    let src_off = d + src_base;
178
179    let x_lo = T::load(lo_ptr.add_offsets(src_off), None, None, &[], None, None, None, false);
180    let x_hi = T::load(hi_ptr.add_offsets(src_off), None, None, &[], None, None, None, false);
181
182    let dst_lo_base: i32 = h * 2 * KEY_DIM;
183    let dst_hi_base: i32 = h * 2 * KEY_DIM + KEY_DIM;
184    let dst_off_lo = (d + dst_lo_base) * (H * W) + (b * c * H * W + n);
185    let dst_off_hi = (d + dst_hi_base) * (H * W) + (b * c * H * W + n);
186
187    T::store(out_ptr.add_offsets(dst_off_lo), x_lo, None, &[], None, None);
188    T::store(out_ptr.add_offsets(dst_off_hi), x_hi, None, &[], None, None);
189}
190
191// ── psa_pack_qkv_backward ─────────────────────────────────────────────────────
192
193/// Backward of `psa_pack_qkv`: scatters `d_packed` back to `d_qkv`.
194///
195/// Grid / block matches the forward pass: `[4 * BH * N, 1, 1]`, block `[KEY_DIM, 1, 1]`.
196///
197/// Uses `atomic_add` because both `psa_pack_qkv_backward` (all four sections)
198/// and `psa_extract_v_backward` (V_lo/V_hi sections) write to the same
199/// `d_qkv` buffer.
200#[kernel]
201pub fn psa_pack_qkv_backward<T: Triton, D: Float, const KEY_DIM: i32>(
202    d_packed_ptr: T::Pointer<D>,  // [4, BH, N, KEY_DIM] gradient of the packed output
203    d_qkv_ptr:   T::Pointer<D>,  // [B, qkv_h, H, W]    gradient accumulation target
204    qkv_h: i32,
205    H: i32,
206    W: i32,
207    B: i32,
208    num_heads: i32,
209) where
210    T::I32Tensor: Tensor<i32, 1>,
211    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
212    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
213{
214    let pid = T::program_id(Axis::X); // [0, 4 * BH * N)
215    let BH = B * num_heads;
216    let N = H * W;
217    let section = pid / (BH * N);
218    let bh = (pid / N) % BH;
219    let n = pid % N;
220    let b = bh / num_heads;
221    let h = bh % num_heads;
222
223    let d = T::arange(0, KEY_DIM);
224
225    // Load from d_packed: flat [4, BH, N, KEY_DIM]
226    let src_base = section * BH * N * KEY_DIM + bh * N * KEY_DIM + n * KEY_DIM;
227    let dx = T::load(d_packed_ptr.add_offsets(d + src_base), None, None, &[], None, None, None, false);
228
229    // atomic_add to d_qkv at the NCHW channel position.
230    let chan_base = h * 4 * KEY_DIM + section * KEY_DIM;
231    let dst_off = (d + chan_base) * (H * W) + (b * qkv_h * H * W + n);
232    T::atomic_add(d_qkv_ptr.add_offsets(dst_off), dx, None, None, None);
233}
234
235// ── psa_extract_v_backward ────────────────────────────────────────────────────
236
237/// Backward of `psa_extract_v_nchw`: scatters `d_v` back to V_lo / V_hi
238/// channels of `d_qkv`.
239///
240/// Grid / block matches the forward: `[BH * N, 1, 1]`, block `[KEY_DIM, 1, 1]`.
241/// Uses `atomic_add` because V_lo / V_hi channels of `d_qkv` are also updated
242/// by `psa_pack_qkv_backward`.
243#[kernel]
244pub fn psa_extract_v_backward<T: Triton, D: Float, const KEY_DIM: i32>(
245    d_v_ptr:   T::Pointer<D>,  // [B, c, H, W]    gradient of the extracted V output
246    d_qkv_ptr: T::Pointer<D>,  // [B, qkv_h, H, W] gradient accumulation target
247    qkv_h: i32,
248    c: i32,
249    H: i32,
250    W: i32,
251    num_heads: i32,
252) where
253    T::I32Tensor: Tensor<i32, 1>,
254    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
255    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
256{
257    let pid = T::program_id(Axis::X); // [0, BH * N)
258    let N = H * W;
259    let bh = pid / N;
260    let n = pid % N;
261    let b = bh / num_heads;
262    let h = bh % num_heads;
263
264    let d = T::arange(0, KEY_DIM);
265
266    // Load from d_v NCHW [B, c, H, W] at V_lo and V_hi channel offsets.
267    let v_lo_src_base = h * 2 * KEY_DIM;
268    let v_hi_src_base = h * 2 * KEY_DIM + KEY_DIM;
269    let src_off_lo = (d + v_lo_src_base) * (H * W) + (b * c * H * W + n);
270    let src_off_hi = (d + v_hi_src_base) * (H * W) + (b * c * H * W + n);
271    let dx_lo = T::load(d_v_ptr.add_offsets(src_off_lo), None, None, &[], None, None, None, false);
272    let dx_hi = T::load(d_v_ptr.add_offsets(src_off_hi), None, None, &[], None, None, None, false);
273
274    // atomic_add to d_qkv at the corresponding QKV channel positions (sections 2 and 3).
275    let qkv_lo_base = h * 4 * KEY_DIM + 2 * KEY_DIM;
276    let qkv_hi_base = h * 4 * KEY_DIM + 3 * KEY_DIM;
277    let dst_off_lo = (d + qkv_lo_base) * (H * W) + (b * qkv_h * H * W + n);
278    let dst_off_hi = (d + qkv_hi_base) * (H * W) + (b * qkv_h * H * W + n);
279    T::atomic_add(d_qkv_ptr.add_offsets(dst_off_lo), dx_lo, None, None, None);
280    T::atomic_add(d_qkv_ptr.add_offsets(dst_off_hi), dx_hi, None, None, None);
281}
282
283// ── psa_merge_attn_backward ───────────────────────────────────────────────────
284
285/// Backward of `psa_merge_attn_nchw`: scatters `d_merged` back to `d_lo` and
286/// `d_hi` flat buffers.
287///
288/// Grid / block matches the forward: `[BH * N, 1, 1]`, block `[KEY_DIM, 1, 1]`.
289/// Regular stores are safe: each `(bh, n, d)` position maps to a unique merged
290/// channel, so `d_lo` and `d_hi` receive no overlapping writes.
291#[kernel]
292pub fn psa_merge_attn_backward<T: Triton, D: Float, const KEY_DIM: i32>(
293    d_merged_ptr: T::Pointer<D>,  // [B, c, H, W]    gradient of the merged output
294    d_lo_ptr:     T::Pointer<D>,  // [BH, N, KEY_DIM] gradient for FA2_lo output
295    d_hi_ptr:     T::Pointer<D>,  // [BH, N, KEY_DIM] gradient for FA2_hi output
296    c: i32,
297    H: i32,
298    W: i32,
299    num_heads: i32,
300) where
301    T::I32Tensor: Tensor<i32, 1>,
302    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
303    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
304{
305    let pid = T::program_id(Axis::X); // [0, BH * N)
306    let N = H * W;
307    let bh = pid / N;
308    let n = pid % N;
309    let b = bh / num_heads;
310    let h = bh % num_heads;
311
312    let d = T::arange(0, KEY_DIM);
313
314    // Load lo and hi slices from d_merged NCHW.
315    let src_lo_base = h * 2 * KEY_DIM;
316    let src_hi_base = h * 2 * KEY_DIM + KEY_DIM;
317    let src_off_lo = (d + src_lo_base) * (H * W) + (b * c * H * W + n);
318    let src_off_hi = (d + src_hi_base) * (H * W) + (b * c * H * W + n);
319    let dx_lo = T::load(d_merged_ptr.add_offsets(src_off_lo), None, None, &[], None, None, None, false);
320    let dx_hi = T::load(d_merged_ptr.add_offsets(src_off_hi), None, None, &[], None, None, None, false);
321
322    // Store to d_lo and d_hi flat [BH, N, KEY_DIM].
323    let dst_base = bh * N * KEY_DIM + n * KEY_DIM;
324    let dst_off  = d + dst_base;
325    T::store(d_lo_ptr.add_offsets(dst_off), dx_lo, None, &[], None, None);
326    T::store(d_hi_ptr.add_offsets(dst_off), dx_hi, None, &[], None, None);
327}
328
329// ── RuntimeOp: PsaPackQkvRuntimeOp ───────────────────────────────────────────
330
331/// Runtime dispatch for the PSA QKV-packing kernel (forward + backward).
332pub struct PsaPackQkvRuntimeOp<D: Float + Send + Sync + 'static> {
333    fwd: PsaPackQkv<D>,
334    bwd: PsaPackQkvBackward<D>,
335    num_heads: usize,
336}
337
338impl<D: Float + Send + Sync + 'static> PsaPackQkvRuntimeOp<D> {
339    /// Builds forward/backward kernels for the given key dimension and head count.
340    pub fn new(key_dim: i32, num_heads: usize) -> Self {
341        Self { fwd: PsaPackQkv::<D>::new(key_dim), bwd: PsaPackQkvBackward::<D>::new(key_dim), num_heads }
342    }
343
344    /// The forward kernel's function name.
345    pub fn kernel_name(&self) -> &str { self.fwd.name }
346    /// The forward kernel's generated source.
347    pub fn forward_source(&self) -> &str { &self.fwd.source }
348    /// The backward kernel's generated source.
349    pub fn backward_source(&self) -> &str { &self.bwd.source }
350}
351
352impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for PsaPackQkvRuntimeOp<D> {
353    fn n_activation_inputs(&self) -> usize { 1 }
354
355    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> { Vec::new() }
356
357    fn pack_args(
358        &self,
359        inputs: &[(teeny_core::model::RawPtr, &[usize])],
360        _params: &[teeny_core::model::RawPtr],
361        output: teeny_core::model::RawPtr,
362        _output_shape: &[usize],
363        _output_row_stride: i32,
364        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
365    ) {
366        // input: [B, qkv_h, H, W]
367        let b = inputs[0].1[0] as i32;
368        let qkv_h = inputs[0].1[1] as i32;
369        let h = inputs[0].1[2] as i32;
370        let w = inputs[0].1[3] as i32;
371        visitor.visit_ptr(inputs[0].0);
372        visitor.visit_ptr(output);
373        visitor.visit_i32(qkv_h);
374        visitor.visit_i32(h);
375        visitor.visit_i32(w);
376        visitor.visit_i32(b);
377        visitor.visit_i32(self.num_heads as i32);
378    }
379
380    fn block(&self) -> [u32; 3] { [128, 1, 1] }
381
382    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
383        // output_shape = [B, 4, num_heads, N, KEY_DIM]
384        // grid = [B * 4 * num_heads * N, 1, 1]
385        [(output_shape[0] * output_shape[1] * output_shape[2] * output_shape[3]) as u32, 1, 1]
386    }
387
388    #[cfg(feature = "training")]
389    fn has_backward(&self) -> bool { true }
390
391    /// kernel args: d_packed, d_qkv, qkv_h, H, W, B, num_heads
392    #[cfg(feature = "training")]
393    fn pack_backward_args(
394        &self,
395        inputs: &[(teeny_core::model::RawPtr, &[usize])],
396        _params: &[teeny_core::model::RawPtr],
397        _output: teeny_core::model::RawPtr,
398        output_shape: &[usize],
399        grad_output: teeny_core::model::RawPtr,
400        _grad_output_row_stride: i32,
401        grad_inputs: &[teeny_core::model::RawPtr],
402        _grad_params: &[teeny_core::model::RawPtr],
403        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
404    ) {
405        // inputs[0].1 = [B, qkv_h, H, W]; output_shape = [4, BH, N, KEY_DIM]
406        let b     = inputs[0].1[0] as i32;
407        let qkv_h = inputs[0].1[1] as i32;
408        let h     = inputs[0].1[2] as i32;
409        let w     = inputs[0].1[3] as i32;
410        let _ = output_shape;
411        visitor.visit_ptr(grad_output);    // d_packed_ptr
412        visitor.visit_ptr(grad_inputs[0]); // d_qkv_ptr
413        visitor.visit_i32(qkv_h);
414        visitor.visit_i32(h);
415        visitor.visit_i32(w);
416        visitor.visit_i32(b);
417        visitor.visit_i32(self.num_heads as i32);
418    }
419
420    #[cfg(feature = "training")]
421    fn backward_block(&self) -> [u32; 3] { [128, 1, 1] }
422
423    /// Grid = `[B * 4 * num_heads * N, 1, 1]` — same layout as the forward pass.
424    #[cfg(feature = "training")]
425    fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
426        [(output_shape[0] * output_shape[1] * output_shape[2] * output_shape[3]) as u32, 1, 1]
427    }
428}
429
430// ── RuntimeOp: PsaExtractVRuntimeOp ──────────────────────────────────────────
431
432/// Runtime dispatch for the PSA V-extraction kernel (forward + backward).
433pub struct PsaExtractVRuntimeOp<D: Float + Send + Sync + 'static> {
434    fwd: PsaExtractVNchw<D>,
435    bwd: PsaExtractVBackward<D>,
436    num_heads: usize,
437}
438
439impl<D: Float + Send + Sync + 'static> PsaExtractVRuntimeOp<D> {
440    /// Builds forward/backward kernels for the given key dimension and head count.
441    pub fn new(key_dim: i32, num_heads: usize) -> Self {
442        Self { fwd: PsaExtractVNchw::<D>::new(key_dim), bwd: PsaExtractVBackward::<D>::new(key_dim), num_heads }
443    }
444
445    /// The forward kernel's function name.
446    pub fn kernel_name(&self) -> &str { self.fwd.name }
447    /// The forward kernel's generated source.
448    pub fn forward_source(&self) -> &str { &self.fwd.source }
449    /// The backward kernel's generated source.
450    pub fn backward_source(&self) -> &str { &self.bwd.source }
451}
452
453impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for PsaExtractVRuntimeOp<D> {
454    fn n_activation_inputs(&self) -> usize { 1 }
455
456    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> { Vec::new() }
457
458    fn pack_args(
459        &self,
460        inputs: &[(teeny_core::model::RawPtr, &[usize])],
461        _params: &[teeny_core::model::RawPtr],
462        output: teeny_core::model::RawPtr,
463        output_shape: &[usize],
464        _output_row_stride: i32,
465        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
466    ) {
467        // input: [B, qkv_h, H, W]; output: [B, c, H, W]
468        let qkv_h = inputs[0].1[1] as i32;
469        let h = inputs[0].1[2] as i32;
470        let w = inputs[0].1[3] as i32;
471        let c = output_shape[1] as i32;
472        visitor.visit_ptr(inputs[0].0);
473        visitor.visit_ptr(output);
474        visitor.visit_i32(qkv_h);
475        visitor.visit_i32(c);
476        visitor.visit_i32(h);
477        visitor.visit_i32(w);
478        visitor.visit_i32(self.num_heads as i32);
479    }
480
481    fn block(&self) -> [u32; 3] { [128, 1, 1] }
482
483    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
484        // output_shape = [B, c, H, W]; grid = [BH * N, 1, 1]
485        let bh = output_shape[0] * self.num_heads;
486        let n = output_shape[2] * output_shape[3];
487        [(bh * n) as u32, 1, 1]
488    }
489
490    #[cfg(feature = "training")]
491    fn has_backward(&self) -> bool { true }
492
493    /// kernel args: d_v, d_qkv, qkv_h, c, H, W, num_heads
494    #[cfg(feature = "training")]
495    fn pack_backward_args(
496        &self,
497        inputs: &[(teeny_core::model::RawPtr, &[usize])],
498        _params: &[teeny_core::model::RawPtr],
499        _output: teeny_core::model::RawPtr,
500        output_shape: &[usize],
501        grad_output: teeny_core::model::RawPtr,
502        _grad_output_row_stride: i32,
503        grad_inputs: &[teeny_core::model::RawPtr],
504        _grad_params: &[teeny_core::model::RawPtr],
505        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
506    ) {
507        // inputs[0].1 = [B, qkv_h, H, W]; output_shape = [B, c, H, W]
508        let qkv_h = inputs[0].1[1] as i32;
509        let h     = inputs[0].1[2] as i32;
510        let w     = inputs[0].1[3] as i32;
511        let c     = output_shape[1] as i32;
512        visitor.visit_ptr(grad_output);    // d_v_ptr
513        visitor.visit_ptr(grad_inputs[0]); // d_qkv_ptr
514        visitor.visit_i32(qkv_h);
515        visitor.visit_i32(c);
516        visitor.visit_i32(h);
517        visitor.visit_i32(w);
518        visitor.visit_i32(self.num_heads as i32);
519    }
520
521    #[cfg(feature = "training")]
522    fn backward_block(&self) -> [u32; 3] { [128, 1, 1] }
523
524    /// Grid = `[BH * N, 1, 1]` — same layout as the forward pass.
525    #[cfg(feature = "training")]
526    fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
527        let bh = output_shape[0] * self.num_heads;
528        let n = output_shape[2] * output_shape[3];
529        [(bh * n) as u32, 1, 1]
530    }
531}
532
533// ── RuntimeOp: PsaMergeAttnRuntimeOp ─────────────────────────────────────────
534
535/// Runtime dispatch for the PSA attention-merge kernel (forward + backward).
536pub struct PsaMergeAttnRuntimeOp<D: Float + Send + Sync + 'static> {
537    fwd: PsaMergeAttnNchw<D>,
538    bwd: PsaMergeAttnBackward<D>,
539    num_heads: usize,
540}
541
542impl<D: Float + Send + Sync + 'static> PsaMergeAttnRuntimeOp<D> {
543    /// Builds forward/backward kernels for the given key dimension and head count.
544    pub fn new(key_dim: i32, num_heads: usize) -> Self {
545        Self { fwd: PsaMergeAttnNchw::<D>::new(key_dim), bwd: PsaMergeAttnBackward::<D>::new(key_dim), num_heads }
546    }
547
548    /// The forward kernel's function name.
549    pub fn kernel_name(&self) -> &str { self.fwd.name }
550    /// The forward kernel's generated source.
551    pub fn forward_source(&self) -> &str { &self.fwd.source }
552    /// The backward kernel's generated source.
553    pub fn backward_source(&self) -> &str { &self.bwd.source }
554}
555
556impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for PsaMergeAttnRuntimeOp<D> {
557    fn n_activation_inputs(&self) -> usize { 2 }
558
559    fn param_shapes(&self, _: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> { Vec::new() }
560
561    fn pack_args(
562        &self,
563        inputs: &[(teeny_core::model::RawPtr, &[usize])],
564        _params: &[teeny_core::model::RawPtr],
565        output: teeny_core::model::RawPtr,
566        output_shape: &[usize],
567        _output_row_stride: i32,
568        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
569    ) {
570        // inputs[0] = o_lo [BH, N, KEY_DIM], inputs[1] = o_hi [BH, N, KEY_DIM]
571        // output_shape = [B, c, H, W]
572        let c = output_shape[1] as i32;
573        let h = output_shape[2] as i32;
574        let w = output_shape[3] as i32;
575        visitor.visit_ptr(inputs[0].0);
576        visitor.visit_ptr(inputs[1].0);
577        visitor.visit_ptr(output);
578        visitor.visit_i32(c);
579        visitor.visit_i32(h);
580        visitor.visit_i32(w);
581        visitor.visit_i32(self.num_heads as i32);
582    }
583
584    fn block(&self) -> [u32; 3] { [128, 1, 1] }
585
586    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
587        // output_shape = [B, c, H, W]; grid = [BH * N, 1, 1]
588        let bh = output_shape[0] * self.num_heads;
589        let n = output_shape[2] * output_shape[3];
590        [(bh * n) as u32, 1, 1]
591    }
592
593    #[cfg(feature = "training")]
594    fn has_backward(&self) -> bool { true }
595
596    /// kernel args: d_merged, d_lo, d_hi, c, H, W, num_heads
597    #[cfg(feature = "training")]
598    fn pack_backward_args(
599        &self,
600        _inputs: &[(teeny_core::model::RawPtr, &[usize])],
601        _params: &[teeny_core::model::RawPtr],
602        _output: teeny_core::model::RawPtr,
603        output_shape: &[usize],
604        grad_output: teeny_core::model::RawPtr,
605        _grad_output_row_stride: i32,
606        grad_inputs: &[teeny_core::model::RawPtr],
607        _grad_params: &[teeny_core::model::RawPtr],
608        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
609    ) {
610        // output_shape = [B, c, H, W]; grad_inputs = [d_lo, d_hi]
611        let c = output_shape[1] as i32;
612        let h = output_shape[2] as i32;
613        let w = output_shape[3] as i32;
614        visitor.visit_ptr(grad_output);    // d_merged_ptr
615        visitor.visit_ptr(grad_inputs[0]); // d_lo_ptr
616        visitor.visit_ptr(grad_inputs[1]); // d_hi_ptr
617        visitor.visit_i32(c);
618        visitor.visit_i32(h);
619        visitor.visit_i32(w);
620        visitor.visit_i32(self.num_heads as i32);
621    }
622
623    #[cfg(feature = "training")]
624    fn backward_block(&self) -> [u32; 3] { [128, 1, 1] }
625
626    /// Grid = `[BH * N, 1, 1]` — same layout as the forward pass.
627    #[cfg(feature = "training")]
628    fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
629        let bh = output_shape[0] * self.num_heads;
630        let n = output_shape[2] * output_shape[3];
631        [(bh * n) as u32, 1, 1]
632    }
633}
634
635// ── psa_fa2_backward ─────────────────────────────────────────────────────────
636
637/// Combined PSA Flash Attention 2 backward pass.
638///
639/// Each CTA handles one `(n, bh)` pair and computes both dQ_n (by iterating
640/// over all K rows) and dK_n / dV_n (by iterating over all Q rows).
641/// This is valid in PSA because `N_CTX_Q == N_CTX_K == N` (self-attention).
642///
643/// **Atomicity**: `dq_ptr` and `dk_ptr` are written with `atomic_add` because
644/// both FA2_lo and FA2_hi backward passes contribute to shared sections 0 and 1
645/// of the d_packed buffer.  `dv_ptr` uses a regular store (each FA2 call owns
646/// an exclusive V section so no overlap exists).
647///
648/// Grid: `(N, BH, 1)` — same shape as the FA2 forward pass.
649#[kernel]
650pub fn psa_fa2_backward<T: Triton, D: Float, const HEAD_DIM: i32>(
651    q_ptr:  T::Pointer<D>,  // [BH, N, HEAD_DIM] Q — section 0 of packed forward buffer
652    k_ptr:  T::Pointer<D>,  // [BH, N, HEAD_DIM] K — section 1 of packed forward buffer
653    v_ptr:  T::Pointer<D>,  // [BH, N, HEAD_DIM] V — section v_section of packed forward buffer
654    o_ptr:  T::Pointer<D>,  // [BH, N, HEAD_DIM] FA2 forward output
655    do_ptr: T::Pointer<D>,  // [BH, N, HEAD_DIM] upstream gradient
656    l_ptr:  T::Pointer<D>,  // [BH * N]          logsumexp saved from forward
657    dq_ptr: T::Pointer<D>,  // [BH, N, HEAD_DIM] atomic_add target for dQ
658    dk_ptr: T::Pointer<D>,  // [BH, N, HEAD_DIM] atomic_add target for dK
659    dv_ptr: T::Pointer<D>,  // [BH, N, HEAD_DIM] store target for dV
660    N: i32,                   // N_CTX (== N_CTX_Q == N_CTX_K in PSA self-attention)
661    softmax_scale: f32,
662) where
663    T::I32Tensor: Tensor<i32, 1>,
664    T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
665    T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
666{
667    let pid_n  = T::program_id(Axis::X); // spatial token [0, N)
668    let pid_bh = T::program_id(Axis::Y); // (batch, head)  [0, BH)
669
670    let row_base = pid_bh * N * HEAD_DIM + pid_n * HEAD_DIM;
671    let l_base   = pid_bh * N + pid_n;
672    let bh_base  = pid_bh * N * HEAD_DIM;
673    let l_bh     = pid_bh * N;
674
675    let d       = T::arange(0, HEAD_DIM);
676    let scale_t = T::full(&[HEAD_DIM], D::from_f64(softmax_scale as f64));
677
678    // Load this row's Q, O, dO and compute D_n = rowsum(O * dO).
679    let q_vec  = T::load(q_ptr.add_offsets(d + row_base),  None, None, &[], None, None, None, false);
680    let o_vec  = T::load(o_ptr.add_offsets(d + row_base),  None, None, &[], None, None, None, false);
681    let do_vec = T::load(do_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
682
683    let d_n = T::sum(o_vec * do_vec, Some(0), false); // scalar
684
685    let l_n_raw = T::load(l_ptr.add_offsets(T::arange(0, 1) + l_base), None, None, &[], None, None, None, false);
686    let l_n     = T::sum(l_n_raw, Some(0), false); // scalar
687
688    // Phase 1: accumulate dQ_n by iterating over all K rows.
689    let mut dq_acc = T::zeros::<D>(&[HEAD_DIM]);
690    for k_row in 0..N {
691        let kv_row_base  = bh_base + k_row * HEAD_DIM;
692        let k_vec        = T::load(k_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
693        let v_vec        = T::load(v_ptr.add_offsets(d + kv_row_base), None, None, &[], None, None, None, false);
694        let qk           = T::sum(q_vec * k_vec, Some(0), false) * scale_t;
695        let p            = T::exp(qk - l_n);
696        let do_dot_v     = T::sum(do_vec * v_vec, Some(0), false);
697        let ds           = p * (do_dot_v - d_n);
698        dq_acc = dq_acc + ds * k_vec * scale_t;
699    }
700    T::atomic_add(dq_ptr.add_offsets(d + row_base), dq_acc, None, None, None);
701
702    // Phase 2: accumulate dK_n and dV_n by iterating over all Q rows.
703    let k_vec_n = T::load(k_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
704    let v_vec_n = T::load(v_ptr.add_offsets(d + row_base), None, None, &[], None, None, None, false);
705    let mut dk_acc = T::zeros::<D>(&[HEAD_DIM]);
706    let mut dv_acc = T::zeros::<D>(&[HEAD_DIM]);
707    for q_row in 0..N {
708        let q_row_base_m  = bh_base + q_row * HEAD_DIM;
709        let l_row_base_m  = l_bh + q_row;
710        let q_vec_m  = T::load(q_ptr.add_offsets(d + q_row_base_m),  None, None, &[], None, None, None, false);
711        let o_vec_m  = T::load(o_ptr.add_offsets(d + q_row_base_m),  None, None, &[], None, None, None, false);
712        let do_vec_m = T::load(do_ptr.add_offsets(d + q_row_base_m), None, None, &[], None, None, None, false);
713        let l_m_raw  = T::load(l_ptr.add_offsets(T::arange(0, 1) + l_row_base_m), None, None, &[], None, None, None, false);
714        let l_m      = T::sum(l_m_raw, Some(0), false);
715        let d_m      = T::sum(o_vec_m * do_vec_m, Some(0), false);
716        let qk       = T::sum(q_vec_m * k_vec_n, Some(0), false) * scale_t;
717        let p        = T::exp(qk - l_m);
718        dv_acc = dv_acc + p * do_vec_m;
719        let do_dot_v_m = T::sum(do_vec_m * v_vec_n, Some(0), false);
720        let ds_m = p * (do_dot_v_m - d_m);
721        dk_acc = dk_acc + ds_m * q_vec_m * scale_t;
722    }
723    T::atomic_add(dk_ptr.add_offsets(d + row_base), dk_acc, None, None, None);
724    T::store(dv_ptr.add_offsets(d + row_base), dv_acc, None, &[], None, None);
725}
726
727// ── CustomOp wrappers ─────────────────────────────────────────────────────────
728//
729// These thin `Arc`-wrapper structs implement `teeny_core::graph::CustomOp` so
730// that `SymTensor::record_custom` can record PSA graph nodes directly in
731// vision-rs without any lowering middleware.  `lower()` hands the pre-built
732// `Arc<RuntimeOp>` straight to `TritonLowering`.
733
734use std::any::Any;
735use std::sync::Arc;
736use teeny_core::{
737    graph::{CustomOp, Shape},
738    model::RuntimeOp,
739};
740
741/// CustomOp for `PsaPackQkvRuntimeOp`.
742///
743/// Graph node: `[B, qkv_h, H, W]` → `[4, BH, N, KEY_DIM]`
744pub struct PsaPackQkvOp<D: Float + Send + Sync + 'static> {
745    inner: Arc<PsaPackQkvRuntimeOp<D>>,
746    num_heads: usize,
747}
748
749impl<D: Float + Send + Sync + 'static> PsaPackQkvOp<D> {
750    /// Creates the graph op for the given key dimension and head count.
751    pub fn new(key_dim: i32, num_heads: usize) -> Self {
752        Self { inner: Arc::new(PsaPackQkvRuntimeOp::<D>::new(key_dim, num_heads)), num_heads }
753    }
754}
755
756impl<D: Float + Send + Sync + 'static> CustomOp for PsaPackQkvOp<D> {
757    fn name(&self) -> &str { "psa_pack_qkv" }
758
759    /// Output shape: `[B, 4, num_heads, N, KEY_DIM]`
760    ///
761    /// B is kept as the leading (potentially-`None`) dimension so that
762    /// `resolve_shape(…, batch_size)` sets it to `batch_size` directly.
763    /// Placing B first avoids the prior bug where `BH = B * num_heads` was
764    /// stored as `None` and resolved to `batch_size` instead of
765    /// `batch_size * num_heads`.
766    fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
767        let s = input_shapes[0];  // [B, qkv_h, H, W]
768        let nh = self.num_heads;
769        vec![
770            s[0],                   // B (dynamic)
771            Some(4),                // sections
772            Some(nh),               // num_heads
773            s[2].and_then(|h| s[3].map(|w| h * w)), // N = H * W
774            s[1].map(|qkv_h| qkv_h / (nh * 4)),    // KEY_DIM
775        ]
776    }
777
778    fn as_any(&self) -> &dyn Any { self }
779
780    fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
781        Some((
782            self.inner.kernel_name().to_string(),
783            self.inner.forward_source().to_string(),
784            "entry_point".to_string(),
785            Arc::clone(&self.inner) as Arc<dyn RuntimeOp>,
786        ))
787    }
788
789    fn lower_backward_source(&self) -> String {
790        self.inner.backward_source().to_string()
791    }
792}
793
794/// CustomOp for `PsaExtractVRuntimeOp`.
795///
796/// Graph node: `[B, qkv_h, H, W]` → `[B, c, H, W]`  (`c = qkv_h / 2`)
797pub struct PsaExtractVOp<D: Float + Send + Sync + 'static>(Arc<PsaExtractVRuntimeOp<D>>);
798
799impl<D: Float + Send + Sync + 'static> PsaExtractVOp<D> {
800    /// Creates the graph op for the given key dimension and head count.
801    pub fn new(key_dim: i32, num_heads: usize) -> Self {
802        Self(Arc::new(PsaExtractVRuntimeOp::<D>::new(key_dim, num_heads)))
803    }
804}
805
806impl<D: Float + Send + Sync + 'static> CustomOp for PsaExtractVOp<D> {
807    fn name(&self) -> &str { "psa_extract_v_nchw" }
808
809    fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
810        let s = input_shapes[0];
811        vec![s[0], s[1].map(|qkv_h| qkv_h / 2), s[2], s[3]]
812    }
813
814    fn as_any(&self) -> &dyn Any { self }
815
816    fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
817        Some((
818            self.0.kernel_name().to_string(),
819            self.0.forward_source().to_string(),
820            "entry_point".to_string(),
821            Arc::clone(&self.0) as Arc<dyn RuntimeOp>,
822        ))
823    }
824
825    fn lower_backward_source(&self) -> String {
826        self.0.backward_source().to_string()
827    }
828}
829
830/// CustomOp for `PsaMergeAttnRuntimeOp`.
831///
832/// Graph node: (`[BH, N, KEY_DIM]`, `[BH, N, KEY_DIM]`) → `[B, c, H, W]`
833///
834/// `h` and `w` must be provided at construction time because they cannot be
835/// derived from `N = H×W` alone at graph-trace time.
836pub struct PsaMergeAttnOp<D: Float + Send + Sync + 'static> {
837    inner: Arc<PsaMergeAttnRuntimeOp<D>>,
838    num_heads: usize,
839    h: usize,
840    w: usize,
841}
842
843impl<D: Float + Send + Sync + 'static> PsaMergeAttnOp<D> {
844    /// Creates the graph op for the given key dimension, head count, and output spatial size.
845    pub fn new(key_dim: i32, num_heads: usize, h: usize, w: usize) -> Self {
846        Self { inner: Arc::new(PsaMergeAttnRuntimeOp::<D>::new(key_dim, num_heads)), num_heads, h, w }
847    }
848}
849
850impl<D: Float + Send + Sync + 'static> CustomOp for PsaMergeAttnOp<D> {
851    fn name(&self) -> &str { "psa_merge_attn_nchw" }
852
853    fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
854        // input: [B, num_heads, N, KEY_DIM]  →  output: [B, c, H, W]
855        let lo = input_shapes[0];
856        let nh = self.num_heads;
857        vec![
858            lo[0],                          // B
859            lo[3].map(|kd| nh * 2 * kd),   // c = num_heads * 2 * KEY_DIM
860            Some(self.h),
861            Some(self.w),
862        ]
863    }
864
865    fn as_any(&self) -> &dyn Any { self }
866
867    fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
868        Some((
869            self.inner.kernel_name().to_string(),
870            self.inner.forward_source().to_string(),
871            "entry_point".to_string(),
872            Arc::clone(&self.inner) as Arc<dyn RuntimeOp>,
873        ))
874    }
875
876    fn lower_backward_source(&self) -> String {
877        self.inner.backward_source().to_string()
878    }
879}
880
881/// CustomOp for `FlashAttn2PsaRuntimeOp` (V_lo or V_hi section).
882///
883/// Graph node: `[4, BH, N, KEY_DIM]` → `[BH, N, KEY_DIM]`
884pub struct FlashAttn2PsaOp<D: Float + Send + Sync + 'static>(Arc<FlashAttn2PsaRuntimeOp<D>>);
885
886impl<D: Float + Send + Sync + 'static> FlashAttn2PsaOp<D> {
887    /// Creates the graph op for attention on the V_lo section.
888    pub fn new_lo(key_dim: i32) -> Self {
889        Self(Arc::new(FlashAttn2PsaRuntimeOp::<D>::new_lo(key_dim)))
890    }
891
892    /// Creates the graph op for attention on the V_hi section.
893    pub fn new_hi(key_dim: i32) -> Self {
894        Self(Arc::new(FlashAttn2PsaRuntimeOp::<D>::new_hi(key_dim)))
895    }
896}
897
898impl<D: Float + Send + Sync + 'static> CustomOp for FlashAttn2PsaOp<D> {
899    fn name(&self) -> &str { "flash_attention2_forward" }
900
901    fn infer_output_shape(&self, input_shapes: &[&Shape]) -> Shape {
902        // input: [B, 4, num_heads, N, KEY_DIM]  →  output: [B, num_heads, N, KEY_DIM]
903        let s = input_shapes[0];
904        vec![s[0], s[2], s[3], s[4]]
905    }
906
907    fn as_any(&self) -> &dyn Any { self }
908
909    fn lower(&self) -> Option<(String, String, String, Arc<dyn RuntimeOp>)> {
910        Some((
911            self.0.kernel_name().to_string(),
912            self.0.forward_source().to_string(),
913            "entry_point".to_string(),
914            Arc::clone(&self.0) as Arc<dyn RuntimeOp>,
915        ))
916    }
917
918    fn lower_backward_source(&self) -> String {
919        self.0.backward_source().to_string()
920    }
921}
922
923// ── RuntimeOp: FlashAttn2PsaRuntimeOp ────────────────────────────────────────
924
925/// RuntimeOp wrapping `flash_attention2_forward` for PSA attention.
926///
927/// Input:  packed QKV buffer, shape `[4, BH, N, KEY_DIM]`
928/// Output: attention result, shape `[BH, N, KEY_DIM]`
929/// Params: `[BH * N]` scratch for FA2 logsumexp `l_ptr`.
930pub struct FlashAttn2PsaRuntimeOp<D: Float + Send + Sync + 'static> {
931    fwd: FlashAttention2Forward<D>,
932    bwd: PsaFa2Backward<D>,
933    v_section: usize,
934}
935
936impl<D: Float + Send + Sync + 'static> FlashAttn2PsaRuntimeOp<D> {
937    /// Attention on V_lo (section index 2).
938    pub fn new_lo(key_dim: i32) -> Self {
939        Self { fwd: FlashAttention2Forward::<D>::new(key_dim), bwd: PsaFa2Backward::<D>::new(key_dim), v_section: 2 }
940    }
941
942    /// Attention on V_hi (section index 3).
943    pub fn new_hi(key_dim: i32) -> Self {
944        Self { fwd: FlashAttention2Forward::<D>::new(key_dim), bwd: PsaFa2Backward::<D>::new(key_dim), v_section: 3 }
945    }
946
947    /// The forward kernel's function name.
948    pub fn kernel_name(&self) -> &str { self.fwd.name }
949    /// The forward kernel's generated source.
950    pub fn forward_source(&self) -> &str { &self.fwd.source }
951    /// The backward kernel's generated source.
952    pub fn backward_source(&self) -> &str { &self.bwd.source }
953}
954
955impl<D: Float + Send + Sync + 'static> teeny_core::model::RuntimeOp for FlashAttn2PsaRuntimeOp<D> {
956    fn n_activation_inputs(&self) -> usize { 1 }
957
958    fn param_shapes(&self, input_shapes: &[&[usize]], _: &[usize]) -> Vec<Vec<usize>> {
959        // input_shapes[0] = [B, 4, num_heads, N, KEY_DIM]
960        let bh = input_shapes[0][0] * input_shapes[0][2]; // B * num_heads
961        let n  = input_shapes[0][3];
962        vec![vec![bh * n]] // l_ptr scratch
963    }
964
965    fn pack_args(
966        &self,
967        inputs: &[(teeny_core::model::RawPtr, &[usize])],
968        params: &[teeny_core::model::RawPtr],
969        output: teeny_core::model::RawPtr,
970        _output_shape: &[usize],
971        _output_row_stride: i32,
972        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
973    ) {
974        // inputs[0].1 = [B, 4, num_heads, N, KEY_DIM]
975        let b  = inputs[0].1[0];
976        let nh = inputs[0].1[2];
977        let n  = inputs[0].1[3];
978        let kd = inputs[0].1[4];
979        let bh = b * nh; // BH = B * num_heads
980        let section_elems = bh * n * kd;
981
982        let base = inputs[0].0 as *mut D;
983        let q_ptr = base as *mut c_void;
984        let k_ptr = unsafe { base.add(section_elems) } as *mut c_void;
985        let v_ptr = unsafe { base.add(self.v_section * section_elems) } as *mut c_void;
986        let softmax_scale = 1.0_f32 / (kd as f32).sqrt();
987
988        visitor.visit_ptr(q_ptr);
989        visitor.visit_ptr(k_ptr);
990        visitor.visit_ptr(v_ptr);
991        visitor.visit_ptr(output);
992        visitor.visit_ptr(params[0]);
993        visitor.visit_i32(n as i32);
994        visitor.visit_i32(n as i32);
995        visitor.visit_f32(softmax_scale);
996        visitor.visit_f32(f32::NEG_INFINITY);
997    }
998
999    fn block(&self) -> [u32; 3] { [1, 1, 1] }
1000
1001    fn grid(&self, output_shape: &[usize]) -> [u32; 3] {
1002        // output_shape = [B, num_heads, N, KEY_DIM]; FA2 grid = (N, BH, 1)
1003        [output_shape[2] as u32, (output_shape[0] * output_shape[1]) as u32, 1]
1004    }
1005
1006    #[cfg(feature = "training")]
1007    fn has_backward(&self) -> bool { true }
1008
1009    /// kernel args: q, k, v, o, do, l, dq, dk, dv, N, softmax_scale
1010    #[cfg(feature = "training")]
1011    fn pack_backward_args(
1012        &self,
1013        inputs: &[(teeny_core::model::RawPtr, &[usize])],
1014        params: &[teeny_core::model::RawPtr],
1015        output: teeny_core::model::RawPtr,
1016        _output_shape: &[usize],
1017        grad_output: teeny_core::model::RawPtr,
1018        _grad_output_row_stride: i32,
1019        grad_inputs: &[teeny_core::model::RawPtr],
1020        _grad_params: &[teeny_core::model::RawPtr],
1021        visitor: &mut dyn teeny_core::device::program::ArgVisitor,
1022    ) {
1023        // inputs[0].1 = [B, 4, num_heads, N, KEY_DIM]
1024        let b  = inputs[0].1[0];
1025        let nh = inputs[0].1[2];
1026        let n  = inputs[0].1[3];
1027        let kd = inputs[0].1[4];
1028        let bh = b * nh;
1029        let section_elems = bh * n * kd;
1030        let softmax_scale = 1.0_f32 / (kd as f32).sqrt();
1031
1032        // Forward Q, K, V pointers from the packed input buffer.
1033        let fwd_base = inputs[0].0 as *mut D;
1034        let q_ptr = fwd_base as *mut c_void;
1035        let k_ptr = unsafe { fwd_base.add(section_elems) } as *mut c_void;
1036        let v_ptr = unsafe { fwd_base.add(self.v_section * section_elems) } as *mut c_void;
1037
1038        // Gradient pointers into d_packed (same layout as packed input).
1039        let d_base = grad_inputs[0] as *mut D;
1040        let dq_ptr = d_base as *mut c_void;
1041        let dk_ptr = unsafe { d_base.add(section_elems) } as *mut c_void;
1042        let dv_ptr = unsafe { d_base.add(self.v_section * section_elems) } as *mut c_void;
1043
1044        visitor.visit_ptr(q_ptr);
1045        visitor.visit_ptr(k_ptr);
1046        visitor.visit_ptr(v_ptr);
1047        visitor.visit_ptr(output);       // o_ptr
1048        visitor.visit_ptr(grad_output);  // do_ptr
1049        visitor.visit_ptr(params[0]);    // l_ptr
1050        visitor.visit_ptr(dq_ptr);
1051        visitor.visit_ptr(dk_ptr);
1052        visitor.visit_ptr(dv_ptr);
1053        visitor.visit_i32(n as i32);     // N
1054        visitor.visit_f32(softmax_scale);
1055    }
1056
1057    #[cfg(feature = "training")]
1058    fn backward_block(&self) -> [u32; 3] { [1, 1, 1] }
1059
1060    /// Grid over `(N, BH, 1)` — same shape as the forward pass.
1061    #[cfg(feature = "training")]
1062    fn backward_grid(&self, _input_shapes: &[&[usize]], output_shape: &[usize]) -> [u32; 3] {
1063        // output_shape = [B, num_heads, N, KEY_DIM]
1064        [output_shape[2] as u32, (output_shape[0] * output_shape[1]) as u32, 1]
1065    }
1066}