#![allow(non_snake_case)]
use teeny_macros::kernel;
use teeny_triton::triton::{
types::{AddOffsets, Comparison},
*,
};
#[kernel]
pub fn cosine_embedding_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
x1_ptr: T::Pointer<f32>,
x2_ptr: T::Pointer<f32>,
y_ptr: T::Pointer<f32>,
out_ptr: T::Pointer<f32>,
_n_rows: i32,
n_dim: i32,
margin: f32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
{
let pid = T::program_id(Axis::X);
let row_base = pid * n_dim;
let col_offs: T::I32Tensor = T::arange(0, BLOCK_SIZE);
let row_offs: T::I32Tensor = col_offs + row_base;
let in_row = col_offs.lt(n_dim);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let x1 = T::load(
x1_ptr.add_offsets(row_offs),
Some(in_row),
Some(zeros),
&[],
None,
None,
None,
false,
);
let x2 = T::load(
x2_ptr.add_offsets(row_offs),
Some(in_row),
Some(zeros),
&[],
None,
None,
None,
false,
);
let dot_raw = T::sum(x1 * x2, Some(0), true);
let sq1_raw = T::sum(x1 * x1, Some(0), true);
let sq2_raw = T::sum(x2 * x2, Some(0), true);
let dot_t = T::zeros::<f32>(&[1]) + dot_raw;
let sq1_t = T::zeros::<f32>(&[1]) + sq1_raw;
let sq2_t = T::zeros::<f32>(&[1]) + sq2_raw;
let norm1 = T::sqrt_rn(sq1_t);
let norm2 = T::sqrt_rn(sq2_t);
let cos_sim = dot_t / (norm1 * norm2);
let y_off: T::I32Tensor = T::arange(0, 1) + pid;
let y: T::Tensor<f32> = T::load(
y_ptr.add_offsets(y_off),
None,
None,
&[],
None,
None,
None,
false,
);
let zeros1 = T::zeros::<f32>(&[1]);
let margin_t = T::full::<f32>(&[1], margin);
let y_is_pos = T::gt(y, zeros1);
let hinge = T::maximum(cos_sim - margin_t, zeros1);
let loss = T::where_(y_is_pos, T::full::<f32>(&[1], 1.0_f32) - cos_sim, hinge);
let out_off: T::I32Tensor = T::arange(0, 1) + pid;
T::store(out_ptr.add_offsets(out_off), loss, None, &[], None, None);
}
#[kernel]
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: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
{
let pid = T::program_id(Axis::X);
let row_base = pid * n_dim;
let col_offs: T::I32Tensor = T::arange(0, BLOCK_SIZE);
let row_offs: T::I32Tensor = col_offs + row_base;
let in_row = col_offs.lt(n_dim);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let x1 = T::load(
x1_ptr.add_offsets(row_offs),
Some(in_row),
Some(zeros),
&[],
None,
None,
None,
false,
);
let x2 = T::load(
x2_ptr.add_offsets(row_offs),
Some(in_row),
Some(zeros),
&[],
None,
None,
None,
false,
);
let dot_t = T::zeros::<f32>(&[1]) + T::sum(x1 * x2, Some(0), true);
let sq1_t = T::zeros::<f32>(&[1]) + T::sum(x1 * x1, Some(0), true);
let sq2_t = T::zeros::<f32>(&[1]) + T::sum(x2 * x2, Some(0), true);
let one = T::full::<f32>(&[1], 1.0_f32);
let inv_norm1 = one / T::sqrt_rn(sq1_t);
let inv_norm2 = one / T::sqrt_rn(sq2_t);
let cos_sim = dot_t * inv_norm1 * inv_norm2;
let scalar_off: T::I32Tensor = T::arange(0, 1) + pid;
let dy: T::Tensor<f32> = T::load(
dy_ptr.add_offsets(scalar_off),
None,
None,
&[],
None,
None,
None,
false,
);
let y: T::Tensor<f32> = T::load(
y_ptr.add_offsets(scalar_off),
None,
None,
&[],
None,
None,
None,
false,
);
let zeros1 = T::zeros::<f32>(&[1]);
let margin_t = T::full::<f32>(&[1], margin);
let y_is_pos = T::gt(y, zeros1);
let cos_gt_margin = T::gt(cos_sim, margin_t);
let neg_dy = T::full::<f32>(&[1], -1.0_f32) * dy;
let coeff = T::where_(y_is_pos, neg_dy, T::where_(cos_gt_margin, dy, zeros1));
let inv_norm1_b = T::broadcast_to(inv_norm1, &[BLOCK_SIZE]);
let inv_norm2_b = T::broadcast_to(inv_norm2, &[BLOCK_SIZE]);
let cos_sim_b = T::broadcast_to(cos_sim, &[BLOCK_SIZE]);
let coeff_b = T::broadcast_to(coeff, &[BLOCK_SIZE]);
let d_cos_dx1 = (x2 * inv_norm2_b - cos_sim_b * x1 * inv_norm1_b) * inv_norm1_b;
let d_cos_dx2 = (x1 * inv_norm1_b - cos_sim_b * x2 * inv_norm2_b) * inv_norm2_b;
let dx1 = coeff_b * d_cos_dx1;
let dx2 = coeff_b * d_cos_dx2;
T::store(
dx1_ptr.add_offsets(row_offs),
dx1,
Some(in_row),
&[],
None,
None,
);
T::store(
dx2_ptr.add_offsets(row_offs),
dx2,
Some(in_row),
&[],
None,
None,
);
}
#[kernel]
pub fn triplet_margin_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
anchor_ptr: T::Pointer<f32>,
positive_ptr: T::Pointer<f32>,
negative_ptr: T::Pointer<f32>,
out_ptr: T::Pointer<f32>,
_n_rows: i32,
n_dim: i32,
margin: f32,
eps: f32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
{
let pid = T::program_id(Axis::X);
let row_base = pid * n_dim;
let col_offs: T::I32Tensor = T::arange(0, BLOCK_SIZE);
let row_offs: T::I32Tensor = col_offs + row_base;
let in_row = col_offs.lt(n_dim);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let a = T::load(
anchor_ptr.add_offsets(row_offs),
Some(in_row),
Some(zeros),
&[],
None,
None,
None,
false,
);
let p = T::load(
positive_ptr.add_offsets(row_offs),
Some(in_row),
Some(zeros),
&[],
None,
None,
None,
false,
);
let n = T::load(
negative_ptr.add_offsets(row_offs),
Some(in_row),
Some(zeros),
&[],
None,
None,
None,
false,
);
let diff_ap = a - p;
let diff_an = a - n;
let eps_t = T::full::<f32>(&[1], eps);
let sq_ap = T::sum(diff_ap * diff_ap, Some(0), true) + eps_t;
let sq_an = T::sum(diff_an * diff_an, Some(0), true) + eps_t;
let d_ap = T::sqrt_rn(sq_ap);
let d_an = T::sqrt_rn(sq_an);
let margin_t = T::full::<f32>(&[1], margin);
let zeros1 = T::zeros::<f32>(&[1]);
let loss = T::maximum(d_ap - d_an + margin_t, zeros1);
let out_off: T::I32Tensor = T::arange(0, 1) + pid;
T::store(out_ptr.add_offsets(out_off), loss, None, &[], None, None);
}
#[kernel]
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: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
{
let pid = T::program_id(Axis::X);
let row_base = pid * n_dim;
let col_offs: T::I32Tensor = T::arange(0, BLOCK_SIZE);
let row_offs: T::I32Tensor = col_offs + row_base;
let in_row = col_offs.lt(n_dim);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let a = T::load(
anchor_ptr.add_offsets(row_offs),
Some(in_row),
Some(zeros),
&[],
None,
None,
None,
false,
);
let p = T::load(
positive_ptr.add_offsets(row_offs),
Some(in_row),
Some(zeros),
&[],
None,
None,
None,
false,
);
let n = T::load(
negative_ptr.add_offsets(row_offs),
Some(in_row),
Some(zeros),
&[],
None,
None,
None,
false,
);
let diff_ap = a - p;
let diff_an = a - n;
let eps_t = T::full::<f32>(&[1], eps);
let sq_ap = T::sum(diff_ap * diff_ap, Some(0), true) + eps_t;
let sq_an = T::sum(diff_an * diff_an, Some(0), true) + eps_t;
let one = T::full::<f32>(&[1], 1.0_f32);
let d_ap = T::sqrt_rn(sq_ap);
let d_an = T::sqrt_rn(sq_an);
let inv_d_ap = one / d_ap;
let inv_d_an = one / d_an;
let margin_t = T::full::<f32>(&[1], margin);
let zeros1 = T::zeros::<f32>(&[1]);
let active = T::gt(d_ap - d_an + margin_t, zeros1);
let scalar_off: T::I32Tensor = T::arange(0, 1) + pid;
let dy: T::Tensor<f32> = T::load(
dy_ptr.add_offsets(scalar_off),
None,
None,
&[],
None,
None,
None,
false,
);
let eff_dy = T::where_(active, dy, zeros1);
let neg_eff_dy = T::full::<f32>(&[1], -1.0_f32) * eff_dy;
let inv_d_ap_b = T::broadcast_to(inv_d_ap, &[BLOCK_SIZE]);
let inv_d_an_b = T::broadcast_to(inv_d_an, &[BLOCK_SIZE]);
let eff_dy_b = T::broadcast_to(eff_dy, &[BLOCK_SIZE]);
let neg_eff_b = T::broadcast_to(neg_eff_dy, &[BLOCK_SIZE]);
let unit_ap = diff_ap * inv_d_ap_b;
let unit_an = diff_an * inv_d_an_b;
let da = eff_dy_b * (unit_ap - unit_an);
let dp = neg_eff_b * unit_ap;
let dn = eff_dy_b * unit_an;
T::store(
da_ptr.add_offsets(row_offs),
da,
Some(in_row),
&[],
None,
None,
);
T::store(
dp_ptr.add_offsets(row_offs),
dp,
Some(in_row),
&[],
None,
None,
);
T::store(
dn_ptr.add_offsets(row_offs),
dn,
Some(in_row),
&[],
None,
None,
);
}