pub fn flash_attention2_backward_dkv<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>,
dk_ptr: T::Pointer<D>,
dv_ptr: T::Pointer<D>,
n_ctx_q: i32,
n_ctx_k: 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
Flash Attention 2 backward: computes dK and dV for one key row.
For each key row n, iterates over all query rows m and accumulates:
dV_n += p_{mn} * dO_m
dK_n += dS_{mn} * Q_m * scale
where p_{mn} = exp(Q_m · K_n * scale − L_m)
dS_{mn} = p_{mn} * (dO_m · V_n − D_m)
D_m = sum(O_m * dO_m)Each CTA owns an exclusive key row so no atomic operations are needed.
Grid: (N_CTX_K, BH, 1).