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])