pub fn psa_fa2_backward<T: Triton, D: Float, const HEAD_DIM: i32>(
q_ptr: T::Pointer<D>,
k_ptr: T::Pointer<D>,
v_ptr: T::Pointer<D>,
o_ptr: T::Pointer<D>,
do_ptr: T::Pointer<D>,
l_ptr: T::Pointer<D>,
dq_ptr: T::Pointer<D>,
dk_ptr: T::Pointer<D>,
dv_ptr: T::Pointer<D>,
N: i32,
softmax_scale: f32,
)Expand description
Combined PSA Flash Attention 2 backward pass.
Each CTA handles one (n, bh) pair and computes both dQ_n (by iterating
over all K rows) and dK_n / dV_n (by iterating over all Q rows).
This is valid in PSA because N_CTX_Q == N_CTX_K == N (self-attention).
Atomicity: dq_ptr and dk_ptr are written with atomic_add because
both FA2_lo and FA2_hi backward passes contribute to shared sections 0 and 1
of the d_packed buffer. dv_ptr uses a regular store (each FA2 call owns
an exclusive V section so no overlap exists).
Grid: (N, BH, 1) — same shape as the FA2 forward pass.