Skip to main content

matmul_backward_da

Function matmul_backward_da 

Source
pub fn matmul_backward_da<T: Triton, D: Num, const BLOCK_M: i32, const BLOCK_N: i32, const BLOCK_K: i32, const GROUP_M: i32>(
    dc_ptr: T::Pointer<D>,
    b_ptr: T::Pointer<D>,
    da_ptr: T::Pointer<D>,
    M: i32,
    N: i32,
    K: i32,
)
Expand description

Backward: dA = dC @ B^T

Grid: one CTA per [BLOCK_M, BLOCK_K] tile of dA. dA[m, k] = sum_n dC[m, n] * B[k, n] (B[k, n] = B^T[n, k])