use crate::{
CubeDevice,
kernel::utils::{address_type, shape_divmod},
};
use crate::{element::CubeElement, tensor::CubeTensor};
use crate::{
kernel::{
AddOp, BitwiseAndOp, BitwiseOrOp, BitwiseXorOp, DivOp, MulOp, PowOp, RemainderOp, SubOp,
launch_binop, launch_binop_int, launch_scalar_binop, launch_scalar_binop_int,
},
ops::max_vector_size,
};
use burn_backend::cubecl::dtype_to_storage_type;
use burn_backend::{DType, Shape, TensorMetadata};
use burn_std::Metadata;
use cubecl::{
calculate_cube_count_elemwise,
ir::{ElemType, dialect::math::IsNanOp},
prelude::*,
std::tensor::layout::linear::LinearViewMut,
};
use cubecl::{client::Client, server::MemoryLayout};
use cubecl::{server::MemoryLayoutDescriptor, std::FastDivmod};
pub fn full<E: CubeElement>(shape: Shape, device: &CubeDevice, value: E) -> CubeTensor {
let client = device.client();
full_client::<E>(client, shape, device.clone(), value)
}
pub fn full_client<E: CubeElement>(
client: Client,
shape: Shape,
device: CubeDevice,
value: E,
) -> CubeTensor {
let dtype = E::dtype();
full_device_dtype(
client,
shape,
device,
InputScalar::new(value, dtype_to_storage_type(dtype)),
dtype,
)
}
pub fn full_device_dtype(
client: Client,
shape: Shape,
device: CubeDevice,
value: InputScalar,
dtype: DType,
) -> CubeTensor {
let empty = empty_device_dtype(client, device, shape, dtype);
fill_device_dtype(empty, value)
}
pub(crate) fn fill_device_dtype(tensor: CubeTensor, value: InputScalar) -> CubeTensor {
#[cube(launch_unchecked, address_type = "dynamic")]
pub fn full_kernel<C: Numeric, N: Size>(
mut tensor: LinearViewMut<'_, Vector<C, N>>,
value: InputScalar,
#[define(C)] _dtype: ElemType,
) {
if !tensor.is_in_bounds(ABSOLUTE_POS) {
terminate!();
}
tensor.write(ABSOLUTE_POS, Vector::new(value.get::<C>()));
}
let num_elems = tensor.meta.num_elements();
let vector_size = max_vector_size(&tensor);
let working_units = num_elems / vector_size as usize;
let cube_dim = CubeDim::new(&tensor.client, working_units);
let cube_count = calculate_cube_count_elemwise(&tensor.client, working_units, cube_dim);
unsafe {
full_kernel::launch_unchecked(
&tensor.client,
cube_count,
cube_dim,
address_type!(tensor),
vector_size,
tensor.clone().into_linear_view(),
value,
dtype_to_storage_type(tensor.dtype),
);
}
tensor
}
pub fn zeros(device: CubeDevice, shape: Shape, dtype: DType) -> CubeTensor {
let client = device.client();
full_device_dtype(
client,
shape,
device,
InputScalar::new(0u32, dtype_to_storage_type(dtype)),
dtype,
)
}
pub fn ones(device: CubeDevice, shape: Shape, dtype: DType) -> CubeTensor {
let client = device.client();
full_device_dtype(
client,
shape,
device,
InputScalar::new(1u32, dtype_to_storage_type(dtype)),
dtype,
)
}
pub fn zeros_client(client: Client, device: CubeDevice, shape: Shape, dtype: DType) -> CubeTensor {
full_device_dtype(
client,
shape,
device,
InputScalar::new(0u32, dtype_to_storage_type(dtype)),
dtype,
)
}
pub fn ones_client(client: Client, device: CubeDevice, shape: Shape, dtype: DType) -> CubeTensor {
full_device_dtype(
client,
shape,
device,
InputScalar::new(1u32, dtype_to_storage_type(dtype)),
dtype,
)
}
pub fn empty_device<E: CubeElement>(
client: Client,
device: CubeDevice,
shape: Shape,
) -> CubeTensor {
let MemoryLayout { memory, strides } = client.empty_tensor(shape.clone(), size_of::<E>());
CubeTensor::new(
client,
memory,
Metadata::new(shape, strides),
device,
E::dtype(),
)
}
pub fn empty_device_dtype(
client: Client,
device: CubeDevice,
shape: Shape,
dtype: DType,
) -> CubeTensor {
let MemoryLayout { memory, strides } = client.empty_tensor(shape.clone(), dtype.size());
CubeTensor::new(client, memory, Metadata::new(shape, strides), device, dtype)
}
pub fn empty_device_contiguous_dtype(
client: Client,
device: CubeDevice,
shape: Shape,
dtype: DType,
) -> CubeTensor {
let descriptor = MemoryLayoutDescriptor::contiguous(shape.clone(), dtype.size());
let MemoryLayout { memory, strides } = client.empty_tensors(vec![descriptor]).remove(0);
CubeTensor::new(client, memory, Metadata::new(shape, strides), device, dtype)
}
pub fn add(lhs: CubeTensor, rhs: CubeTensor) -> CubeTensor {
launch_binop::<AddOp>(lhs, rhs)
}
pub fn add_scalar(lhs: CubeTensor, rhs: InputScalar) -> CubeTensor {
launch_scalar_binop::<AddOp>(lhs, rhs)
}
pub fn sub(lhs: CubeTensor, rhs: CubeTensor) -> CubeTensor {
launch_binop::<SubOp>(lhs, rhs)
}
pub fn sub_scalar(lhs: CubeTensor, rhs: InputScalar) -> CubeTensor {
launch_scalar_binop::<SubOp>(lhs, rhs)
}
pub fn mul(lhs: CubeTensor, rhs: CubeTensor) -> CubeTensor {
launch_binop::<MulOp>(lhs, rhs)
}
pub fn mul_scalar(lhs: CubeTensor, rhs: InputScalar) -> CubeTensor {
launch_scalar_binop::<MulOp>(lhs, rhs)
}
pub fn div(lhs: CubeTensor, rhs: CubeTensor) -> CubeTensor {
launch_binop::<DivOp>(lhs, rhs)
}
pub fn div_scalar(lhs: CubeTensor, rhs: InputScalar) -> CubeTensor {
launch_scalar_binop::<DivOp>(lhs, rhs)
}
pub fn remainder(lhs: CubeTensor, rhs: CubeTensor) -> CubeTensor {
launch_binop::<RemainderOp>(lhs, rhs)
}
pub fn remainder_scalar(lhs: CubeTensor, rhs: InputScalar) -> CubeTensor {
launch_scalar_binop::<RemainderOp>(lhs, rhs)
}
pub fn pow(lhs: CubeTensor, rhs: CubeTensor) -> CubeTensor {
launch_binop::<PowOp>(lhs, rhs)
}
pub fn bitwise_and(lhs: CubeTensor, rhs: CubeTensor) -> CubeTensor {
launch_binop_int::<BitwiseAndOp>(lhs, rhs)
}
pub fn bitwise_and_scalar(lhs: CubeTensor, rhs: InputScalar) -> CubeTensor {
launch_scalar_binop_int::<BitwiseAndOp>(lhs, rhs)
}
pub fn bitwise_or(lhs: CubeTensor, rhs: CubeTensor) -> CubeTensor {
launch_binop_int::<BitwiseOrOp>(lhs, rhs)
}
pub fn bitwise_or_scalar(lhs: CubeTensor, rhs: InputScalar) -> CubeTensor {
launch_scalar_binop_int::<BitwiseOrOp>(lhs, rhs)
}
pub fn bitwise_xor(lhs: CubeTensor, rhs: CubeTensor) -> CubeTensor {
launch_binop_int::<BitwiseXorOp>(lhs, rhs)
}
pub fn bitwise_xor_scalar(lhs: CubeTensor, rhs: InputScalar) -> CubeTensor {
launch_scalar_binop_int::<BitwiseXorOp>(lhs, rhs)
}
pub(crate) trait CumulativeOpFamily: Send + Sync + 'static {
type CumulativeOp<C: Numeric>: CumulativeOp<C>;
}
#[cube]
pub(crate) trait CumulativeOp<C: Numeric>: 'static + Send + Sync {
fn execute(lhs: C, rhs: C) -> C;
fn init_value(first_element: C) -> C;
}
struct SumOp;
struct ProdOp;
struct MaxOp;
struct MinOp;
#[cube]
fn numeric_is_nan<N: Numeric>(value: N) -> bool {
intrinsic!(|scope| {
let is_nan = IsNanOp::new(scope.ctx_mut(), value.read_value(scope));
scope.register_with_result(&is_nan).into()
})
}
#[cube]
fn cumulative_max<N: Numeric>(lhs: N, rhs: N) -> N {
let elem_type = elem_type_of::<N>();
if comptime!(elem_type.is_float()) {
if numeric_is_nan::<N>(lhs) {
lhs
} else if numeric_is_nan::<N>(rhs) {
rhs
} else {
max(lhs, rhs)
}
} else {
max(lhs, rhs)
}
}
#[cube]
fn cumulative_min<N: Numeric>(lhs: N, rhs: N) -> N {
let elem_type = elem_type_of::<N>();
if comptime!(elem_type.is_float()) {
if numeric_is_nan::<N>(lhs) {
lhs
} else if numeric_is_nan::<N>(rhs) {
rhs
} else {
min(lhs, rhs)
}
} else {
min(lhs, rhs)
}
}
impl CumulativeOpFamily for SumOp {
type CumulativeOp<C: Numeric> = Self;
}
impl CumulativeOpFamily for ProdOp {
type CumulativeOp<C: Numeric> = Self;
}
impl CumulativeOpFamily for MaxOp {
type CumulativeOp<C: Numeric> = Self;
}
impl CumulativeOpFamily for MinOp {
type CumulativeOp<C: Numeric> = Self;
}
#[cube]
impl<N: Numeric> CumulativeOp<N> for SumOp {
fn execute(lhs: N, rhs: N) -> N {
lhs + rhs
}
fn init_value(_first_element: N) -> N {
N::zero()
}
}
#[cube]
impl<N: Numeric> CumulativeOp<N> for ProdOp {
fn execute(lhs: N, rhs: N) -> N {
lhs * rhs
}
fn init_value(_first_element: N) -> N {
N::from_int(1)
}
}
#[cube]
impl<N: Numeric> CumulativeOp<N> for MaxOp {
fn execute(lhs: N, rhs: N) -> N {
cumulative_max::<N>(lhs, rhs)
}
fn init_value(first_element: N) -> N {
first_element
}
}
#[cube]
impl<N: Numeric> CumulativeOp<N> for MinOp {
fn execute(lhs: N, rhs: N) -> N {
cumulative_min::<N>(lhs, rhs)
}
fn init_value(first_element: N) -> N {
first_element
}
}
#[cube(launch_unchecked, address_type = "dynamic")]
fn cumulative_kernel<C: Numeric, O: CumulativeOpFamily>(
input: &Tensor<C>,
mut output: LinearViewMut<'_, C>,
shape: Sequence<FastDivmod<usize>>,
#[comptime] dim: usize,
#[define(C)] _dtype: ElemType,
) {
if !output.is_in_bounds(ABSOLUTE_POS) {
terminate!();
}
let rank = comptime![shape.len()];
let dim_stride = input.stride(dim);
let mut remainder = ABSOLUTE_POS;
let mut offset = 0;
let mut dim_idx = 0;
#[unroll]
for i in 0..shape.len() {
let i = comptime![rank - i - 1];
let (rem, local_idx) = shape.index(i).div_mod(remainder);
remainder = rem;
if i == dim {
dim_idx = local_idx;
} else {
offset += local_idx * input.stride(i);
}
}
let first_read_idx = offset + dim_idx * dim_stride;
let first_elem = input[first_read_idx];
let mut result = O::CumulativeOp::<C>::init_value(first_elem);
for i in 0..=dim_idx {
let read_idx = offset + i * dim_stride;
result = O::CumulativeOp::<C>::execute(result, input[read_idx]);
}
output.write(ABSOLUTE_POS, result);
}
pub fn cumsum(input: CubeTensor, dim: usize) -> CubeTensor {
cumulative_op::<SumOp>(input, dim)
}
pub fn cumprod(input: CubeTensor, dim: usize) -> CubeTensor {
cumulative_op::<ProdOp>(input, dim)
}
pub fn cummin(input: CubeTensor, dim: usize) -> CubeTensor {
cumulative_op::<MinOp>(input, dim)
}
pub fn cummax(input: CubeTensor, dim: usize) -> CubeTensor {
cumulative_op::<MaxOp>(input, dim)
}
fn cumulative_op<O: CumulativeOpFamily>(input: CubeTensor, dim: usize) -> CubeTensor {
let client = input.client.clone();
let device = input.device.clone();
let output = empty_device_dtype(client.clone(), device, input.shape(), input.dtype);
let num_elems = output.meta.num_elements();
let working_units = num_elems;
let cube_dim = CubeDim::new(&client, working_units);
let cube_count = calculate_cube_count_elemwise(&client, working_units, cube_dim);
let shape = shape_divmod(&input);
unsafe {
cumulative_kernel::launch_unchecked::<O>(
&client,
cube_count,
cube_dim,
address_type!(input, output),
input.into_tensor_arg(),
output.clone().into_linear_view(),
shape,
dim,
dtype_to_storage_type(output.dtype),
);
}
output
}