Skip to main content

triplet_margin_loss_backward

Function triplet_margin_loss_backward 

Source
pub fn triplet_margin_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
    dy_ptr: T::Pointer<f32>,
    anchor_ptr: T::Pointer<f32>,
    positive_ptr: T::Pointer<f32>,
    negative_ptr: T::Pointer<f32>,
    da_ptr: T::Pointer<f32>,
    dp_ptr: T::Pointer<f32>,
    dn_ptr: T::Pointer<f32>,
    _n_rows: i32,
    n_dim: i32,
    margin: f32,
    eps: 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

Triplet margin loss backward (per-row).

When active (d(a,p) - d(a,n) + margin > 0):

da[k] = dy * ((a[k]-p[k])/d(a,p) - (a[k]-n[k])/d(a,n))
dp[k] = dy * (p[k]-a[k])/d(a,p)
dn[k] = dy * (a[k]-n[k])/d(a,n)

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