pub fn cosine_embedding_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
dy_ptr: T::Pointer<f32>,
x1_ptr: T::Pointer<f32>,
x2_ptr: T::Pointer<f32>,
y_ptr: T::Pointer<f32>,
dx1_ptr: T::Pointer<f32>,
dx2_ptr: T::Pointer<f32>,
_n_rows: i32,
n_dim: i32,
margin: f32,
)where
T::I32Tensor: Tensor<i32, 1> + Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,Expand description
Cosine embedding loss backward (per-row).
Let c = cos_sim, n1 = ||x1||, n2 = ||x2||, r1 = 1/n1, r2 = 1/n2.
dc/dx1[k] = (x2[k]*r2 - c*x1[k]*r1) * r1
dc/dx2[k] = (x1[k]*r1 - c*x2[k]*r2) * r2
coeff = -dy if y == 1
= dy if y == -1 and cos_sim > margin
= 0 otherwise
dx1 = coeff * dc/dx1, dx2 = coeff * dc/dx2Grid: [n_rows, 1, 1]. BLOCK_SIZE must equal next_power_of_two(n_dim).