Skip to main content

flash_attention2_backward_dkv

Function flash_attention2_backward_dkv 

Source
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).