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