#![allow(non_snake_case)]
use teeny_macros::kernel;
use teeny_triton::triton::{
types::{AddOffsets, Comparison},
*,
};
#[kernel]
pub fn bce_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
input_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
out_ptr: T::Pointer<f32>,
n_elements: i32,
) 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let inp = T::load(
input_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let one = T::full(&[BLOCK_SIZE], 1.0_f32);
let eps = T::full(&[BLOCK_SIZE], 1e-7_f32);
let one_minus_eps = T::full(&[BLOCK_SIZE], 1.0_f32 - 1e-7_f32);
let inp_c = T::clamp(inp, eps, one_minus_eps);
let loss = T::full(&[BLOCK_SIZE], -1.0_f32)
* (tgt * T::log(inp_c) + (one - tgt) * T::log(one - inp_c));
T::store(
out_ptr.add_offsets(offsets),
loss,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn bce_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
dy_ptr: T::Pointer<f32>,
input_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
dx_ptr: T::Pointer<f32>,
n_elements: i32,
) 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let dy = T::load(
dy_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let inp = T::load(
input_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let one = T::full(&[BLOCK_SIZE], 1.0_f32);
let eps = T::full(&[BLOCK_SIZE], 1e-7_f32);
let one_minus_eps = T::full(&[BLOCK_SIZE], 1.0_f32 - 1e-7_f32);
let inp_c = T::clamp(inp, eps, one_minus_eps);
let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
let dx_raw = neg_one * (tgt / inp_c - (one - tgt) / (one - inp_c));
let dx = dx_raw * dy;
T::store(
dx_ptr.add_offsets(offsets),
dx,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn bce_with_logits_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
input_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
out_ptr: T::Pointer<f32>,
n_elements: i32,
) 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let inp = T::load(
input_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let one = T::full(&[BLOCK_SIZE], 1.0_f32);
let relu_x = T::maximum(inp, zeros);
let neg_abs_x = T::full(&[BLOCK_SIZE], -1.0_f32) * T::abs(inp);
let loss = relu_x - inp * tgt + T::log(one + T::exp(neg_abs_x));
T::store(
out_ptr.add_offsets(offsets),
loss,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn bce_with_logits_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
dy_ptr: T::Pointer<f32>,
input_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
dx_ptr: T::Pointer<f32>,
n_elements: i32,
) 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let dy = T::load(
dy_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let inp = T::load(
input_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let one = T::full(&[BLOCK_SIZE], 1.0_f32);
let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
let sig = one / (one + T::exp(neg_one * inp));
let dx = (sig - tgt) * dy;
T::store(
dx_ptr.add_offsets(offsets),
dx,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn soft_margin_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
input_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
out_ptr: T::Pointer<f32>,
n_elements: i32,
) 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let inp = T::load(
input_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let one = T::full(&[BLOCK_SIZE], 1.0_f32);
let neg_tx = T::full(&[BLOCK_SIZE], -1.0_f32) * tgt * inp;
let loss = T::maximum(neg_tx, zeros)
+ T::log(one + T::exp(T::full(&[BLOCK_SIZE], -1.0_f32) * T::abs(tgt * inp)));
T::store(
out_ptr.add_offsets(offsets),
loss,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn soft_margin_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
dy_ptr: T::Pointer<f32>,
input_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
dx_ptr: T::Pointer<f32>,
n_elements: i32,
) 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let dy = T::load(
dy_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let inp = T::load(
input_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
let one = T::full(&[BLOCK_SIZE], 1.0_f32);
let neg_tx = neg_one * tgt * inp;
let sig_neg_tx = one / (one + T::exp(neg_one * neg_tx));
let dx = neg_one * tgt * sig_neg_tx * dy;
T::store(
dx_ptr.add_offsets(offsets),
dx,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn kl_div_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
input_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
out_ptr: T::Pointer<f32>,
n_elements: i32,
) 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let inp = T::load(
input_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let eps = T::full(&[BLOCK_SIZE], 1e-10_f32);
let tgt_safe = T::maximum(tgt, eps);
let loss_raw = tgt * (T::log(tgt_safe) - inp);
let positive = T::gt(tgt, zeros);
let loss = T::where_(positive, loss_raw, zeros);
T::store(
out_ptr.add_offsets(offsets),
loss,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn kl_div_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
dy_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
dx_ptr: T::Pointer<f32>,
n_elements: i32,
) 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let dy = T::load(
dy_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let neg_one = T::full(&[BLOCK_SIZE], -1.0_f32);
let dx = neg_one * tgt * dy;
T::store(
dx_ptr.add_offsets(offsets),
dx,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn poisson_nll_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
input_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
out_ptr: T::Pointer<f32>,
n_elements: i32,
) 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let inp = T::load(
input_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let loss = T::exp(inp) - tgt * inp;
T::store(
out_ptr.add_offsets(offsets),
loss,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn poisson_nll_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
dy_ptr: T::Pointer<f32>,
input_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
dx_ptr: T::Pointer<f32>,
n_elements: i32,
) 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let dy = T::load(
dy_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let inp = T::load(
input_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let dx = (T::exp(inp) - tgt) * dy;
T::store(
dx_ptr.add_offsets(offsets),
dx,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn gaussian_nll_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
input_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
var_ptr: T::Pointer<f32>,
out_ptr: T::Pointer<f32>,
n_elements: i32,
eps_var: 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let inp = T::load(
input_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let var = T::load(
var_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let eps_t = T::full(&[BLOCK_SIZE], eps_var);
let half = T::full(&[BLOCK_SIZE], 0.5_f32);
let var_c = T::maximum(var, eps_t);
let diff = inp - tgt;
let loss = half * (T::log(var_c) + diff * diff / var_c);
T::store(
out_ptr.add_offsets(offsets),
loss,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn gaussian_nll_loss_backward_input<T: Triton, const BLOCK_SIZE: i32>(
dy_ptr: T::Pointer<f32>,
input_ptr: T::Pointer<f32>,
target_ptr: T::Pointer<f32>,
var_ptr: T::Pointer<f32>,
dx_ptr: T::Pointer<f32>,
n_elements: i32,
eps_var: 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 block_start = pid * BLOCK_SIZE;
let offsets = T::arange(0, BLOCK_SIZE) + block_start;
let in_bounds = offsets.lt(n_elements);
let zeros = T::zeros::<f32>(&[BLOCK_SIZE]);
let dy = T::load(
dy_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let inp = T::load(
input_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let tgt = T::load(
target_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let var = T::load(
var_ptr.add_offsets(offsets),
Some(in_bounds),
Some(zeros),
&[],
None,
None,
None,
false,
);
let eps_t = T::full(&[BLOCK_SIZE], eps_var);
let var_c = T::maximum(var, eps_t);
let dx = (inp - tgt) / var_c * dy;
T::store(
dx_ptr.add_offsets(offsets),
dx,
Some(in_bounds),
&[],
None,
None,
);
}