#![allow(non_snake_case)]
use teeny_macros::kernel;
use teeny_triton::triton::{
types::{AddOffsets, Comparison},
*,
};
#[kernel]
pub fn nll_loss_forward<T: Triton>(
log_probs_ptr: T::Pointer<f32>,
targets_ptr: T::Pointer<i32>,
out_ptr: T::Pointer<f32>,
_n_rows: i32,
n_cols: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<i32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<i32>>>,
T::Pointer<f32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<f32>>>,
T::Tensor<i32>: types::Tensor<i32, 1>,
T::Pointer<f32>: AddOffsets<i32, 1, T::Tensor<i32>, Output = T::Tensor<T::Pointer<f32>>>,
{
let pid = T::program_id(Axis::X);
let tgt_off: T::I32Tensor = T::arange(0, 1) + pid;
let tgt: T::Tensor<i32> = T::load(
targets_ptr.add_offsets(tgt_off),
None,
None,
&[],
None,
None,
None,
false,
);
let base: T::Tensor<i32> = T::full::<i32>(&[1], pid * n_cols);
let flat_off: T::Tensor<i32> = base + tgt;
let lp: T::Tensor<f32> = T::load(
log_probs_ptr.add_offsets(flat_off),
None,
None,
&[],
None,
None,
None,
false,
);
let loss = T::full(&[1], -1.0_f32) * lp;
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 nll_loss_backward<T: Triton>(
dy_ptr: T::Pointer<f32>,
targets_ptr: T::Pointer<i32>,
dx_ptr: T::Pointer<f32>,
_n_rows: i32,
n_cols: 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>>>,
T::Pointer<i32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<i32>>>,
T::Tensor<i32>: types::Tensor<i32, 1>,
T::Pointer<f32>: AddOffsets<i32, 1, T::Tensor<i32>, Output = T::Tensor<T::Pointer<f32>>>,
{
let pid = T::program_id(Axis::X);
let dy_off: T::I32Tensor = T::arange(0, 1) + pid;
let dy: T::Tensor<f32> = T::load(
dy_ptr.add_offsets(dy_off),
None,
None,
&[],
None,
None,
None,
false,
);
let tgt_off: T::I32Tensor = T::arange(0, 1) + pid;
let tgt: T::Tensor<i32> = T::load(
targets_ptr.add_offsets(tgt_off),
None,
None,
&[],
None,
None,
None,
false,
);
let base: T::Tensor<i32> = T::full::<i32>(&[1], pid * n_cols);
let flat_off: T::Tensor<i32> = base + tgt;
let neg_dy = T::full(&[1], -1.0_f32) * dy;
T::store(dx_ptr.add_offsets(flat_off), neg_dy, None, &[], None, None);
}
#[kernel]
pub fn cross_entropy_loss_forward<T: Triton, const BLOCK_SIZE: i32>(
input_ptr: T::Pointer<f32>,
targets_ptr: T::Pointer<i32>,
out_ptr: T::Pointer<f32>,
_n_rows: i32,
n_cols: 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>>>,
T::Pointer<i32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<i32>>>,
T::Tensor<i32>: types::Tensor<i32, 1>,
T::Pointer<f32>: AddOffsets<i32, 1, T::Tensor<i32>, Output = T::Tensor<T::Pointer<f32>>>,
{
let pid = T::program_id(Axis::X);
let row_base = pid * n_cols;
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_cols);
let neg_inf = T::full(&[BLOCK_SIZE], -3.4028235e38_f32);
let row = T::load(
input_ptr.add_offsets(row_offs),
Some(in_row),
Some(neg_inf),
&[],
None,
None,
None,
false,
);
let row_max = T::max(row, Some(0), true); let row_shifted = row - row_max; let exp_row = T::exp(row_shifted);
let sum_exp = T::sum(exp_row, Some(0), true); let log_sum_exp = T::log(sum_exp) + row_max;
let tgt_off: T::I32Tensor = T::arange(0, 1) + pid;
let tgt: T::Tensor<i32> = T::load(
targets_ptr.add_offsets(tgt_off),
None,
None,
&[],
None,
None,
None,
false,
);
let base: T::Tensor<i32> = T::full::<i32>(&[1], row_base);
let flat_off: T::Tensor<i32> = base + tgt;
let x_target: T::Tensor<f32> = T::load(
input_ptr.add_offsets(flat_off),
None,
None,
&[],
None,
None,
None,
false,
);
let loss = log_sum_exp - x_target;
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 cross_entropy_loss_backward<T: Triton, const BLOCK_SIZE: i32>(
dy_ptr: T::Pointer<f32>,
input_ptr: T::Pointer<f32>,
targets_ptr: T::Pointer<i32>,
dx_ptr: T::Pointer<f32>,
_n_rows: i32,
n_cols: 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>>>,
T::Pointer<i32>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<i32>>>,
T::Tensor<i32>: types::Tensor<i32, 1>,
T::Pointer<f32>: AddOffsets<i32, 1, T::Tensor<i32>, Output = T::Tensor<T::Pointer<f32>>>,
{
let pid = T::program_id(Axis::X);
let row_base = pid * n_cols;
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_cols);
let dy_off: T::I32Tensor = T::arange(0, 1) + pid;
let dy: T::Tensor<f32> = T::load(
dy_ptr.add_offsets(dy_off),
None,
None,
&[],
None,
None,
None,
false,
);
let neg_inf = T::full(&[BLOCK_SIZE], -3.4028235e38_f32);
let row = T::load(
input_ptr.add_offsets(row_offs),
Some(in_row),
Some(neg_inf),
&[],
None,
None,
None,
false,
);
let sm = T::softmax(row, None, false, false);
let tgt_off: T::I32Tensor = T::arange(0, 1) + pid;
let tgt: T::Tensor<i32> = T::load(
targets_ptr.add_offsets(tgt_off),
None,
None,
&[],
None,
None,
None,
false,
);
let dy_bcast = T::broadcast_to(dy, &[BLOCK_SIZE]);
let dx_row = dy_bcast * sm;
T::store(
dx_ptr.add_offsets(row_offs),
dx_row,
Some(in_row),
&[],
None,
None,
);
let base: T::Tensor<i32> = T::full::<i32>(&[1], row_base);
let flat_off: T::Tensor<i32> = base + tgt;
let neg_dy = T::full(&[1], -1.0_f32) * dy;
T::atomic_add(dx_ptr.add_offsets(flat_off), neg_dy, None, None, None);
}
#[kernel]
pub fn multilabel_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_one = T::full(&[BLOCK_SIZE], -1.0_f32);
let relu_x = T::maximum(inp, zeros);
let neg_abs_x = neg_one * 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 multilabel_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 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,
);
}