#![allow(non_snake_case)]
use teeny_core::dtype::Float;
use teeny_macros::kernel;
use teeny_triton::triton::{
types::{AddOffsets, Comparison},
*,
};
#[kernel(backward = TanhBackward)]
pub fn tanh_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_elements: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
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 x = T::load(
x_ptr.add_offsets(offsets),
Some(in_bounds),
None,
&[],
None,
None,
None,
false,
);
let one = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
let two = T::full(&[BLOCK_SIZE], D::from_f64(2.0));
let neg2 = T::full(&[BLOCK_SIZE], D::from_f64(-2.0));
let s2x = one / (one + T::exp(neg2 * x));
let y = two * s2x - one;
T::store(
y_ptr.add_offsets(offsets),
y,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn tanh_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
dy_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
dx_ptr: T::Pointer<D>,
n_elements: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
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 dy = T::load(
dy_ptr.add_offsets(offsets),
Some(in_bounds),
None,
&[],
None,
None,
None,
false,
);
let y = T::load(
y_ptr.add_offsets(offsets),
Some(in_bounds),
None,
&[],
None,
None,
None,
false,
);
let one = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
let dx = dy * (one - y * y);
T::store(
dx_ptr.add_offsets(offsets),
dx,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel(backward = TanhshrinkBackward)]
pub fn tanhshrink_forward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
n_elements: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
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 x = T::load(
x_ptr.add_offsets(offsets),
Some(in_bounds),
None,
&[],
None,
None,
None,
false,
);
let one = T::full(&[BLOCK_SIZE], D::from_f64(1.0));
let two = T::full(&[BLOCK_SIZE], D::from_f64(2.0));
let neg2 = T::full(&[BLOCK_SIZE], D::from_f64(-2.0));
let s2x = one / (one + T::exp(neg2 * x));
let tanh_x = two * s2x - one;
let y = x - tanh_x;
T::store(
y_ptr.add_offsets(offsets),
y,
Some(in_bounds),
&[],
None,
None,
);
}
#[kernel]
pub fn tanhshrink_backward<T: Triton, D: Float, const BLOCK_SIZE: i32>(
dy_ptr: T::Pointer<D>,
x_ptr: T::Pointer<D>,
y_ptr: T::Pointer<D>,
dx_ptr: T::Pointer<D>,
n_elements: i32,
) where
T::I32Tensor: types::Tensor<i32, 1>,
T::I32Tensor: Comparison<i32, BoolTensor = T::BoolTensor>,
T::Pointer<D>: AddOffsets<i32, 1, T::I32Tensor, Output = T::Tensor<T::Pointer<D>>>,
{
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 dy = T::load(
dy_ptr.add_offsets(offsets),
Some(in_bounds),
None,
&[],
None,
None,
None,
false,
);
let x = T::load(
x_ptr.add_offsets(offsets),
Some(in_bounds),
None,
&[],
None,
None,
None,
false,
);
let y = T::load(
y_ptr.add_offsets(offsets),
Some(in_bounds),
None,
&[],
None,
None,
None,
false,
);
let tanh_x = x - y;
let dx = dy * tanh_x * tanh_x;
T::store(
dx_ptr.add_offsets(offsets),
dx,
Some(in_bounds),
&[],
None,
None,
);
}
pub struct TanhOp<D: Float> {
pub forward: TanhForward<D>,
pub backward: TanhBackward<D>,
}
pub struct TanhshrinkOp<D: Float> {
pub forward: TanhshrinkForward<D>,
pub backward: TanhshrinkBackward<D>,
}