Skip to main content

cosine_embedding_loss_backward

Function cosine_embedding_loss_backward 

Source
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/dx2

Grid: [n_rows, 1, 1]. BLOCK_SIZE must equal next_power_of_two(n_dim).