Skip to main content

psa_fa2_backward

Function psa_fa2_backward 

Source
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,
)
where T::I32Tensor: Tensor<i32, 1> + Comparison<i32, BoolTensor = T::BoolTensor>, T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
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.