use crate::{
ir::{Metadata, Scope},
prelude::*,
unexpanded,
};
use alloc::boxed::Box;
use core::ops::{Deref, DerefMut};
use cubecl_ir::VectorSize;
use crate as cubecl;
#[derive(CubeType)]
pub struct Tensor<T: CubePrimitive> {
pub(super) meta: TensorMeta,
pub(super) buffer: [T],
}
#[derive(CubeType, Clone)]
#[expand(derive(Clone))]
pub struct OwnedTensor<T: CubePrimitive> {
#[allow(unused)]
pub(super) meta: TensorMeta,
pub(super) buffer: Box<[T]>,
}
impl<T: CubePrimitive> TensorExpand<T> {
pub fn __expand_from_parts(meta: TensorMetaExpand, buffer: NativeExpand<[T]>) -> Self {
Self { meta, buffer }
}
}
#[cube]
impl<T: CubePrimitive> OwnedTensor<T> {
pub fn from_parts(meta: TensorMeta, buffer: Box<[T]>) -> Self {
OwnedTensor::<T> { meta, buffer }
}
}
#[cube]
impl<T: CubePrimitive> OwnedTensor<T> {
pub fn as_slice(&self) -> &[T] {
&self.buffer
}
pub fn as_mut_slice(&mut self) -> &mut [T] {
&mut self.buffer
}
}
#[cube]
impl<T: CubePrimitive> Tensor<T> {
pub fn as_slice(&self) -> &[T] {
&self.buffer
}
pub fn as_mut_slice(&mut self) -> &mut [T] {
&mut self.buffer
}
}
mod metadata {
use cubecl_ir::Value;
use super::*;
use crate::ir::{Arithmetic, BinaryOperands, Instruction};
#[cube]
impl<T: CubePrimitive> Tensor<T> {
pub fn stride(&self, dim: usize) -> usize {
intrinsic!(|scope| {
let dim: Value = dim.into();
let list = self.__extract_list(scope);
let out = scope.create_value(usize::__expand_as_type(scope));
scope.register(Instruction::new(Metadata::Stride { dim, list }, out));
out.into()
})
}
pub fn shape(&self, dim: usize) -> usize {
intrinsic!(|scope| {
let dim: Value = dim.into();
let list = self.__extract_list(scope);
let out = scope.create_value(usize::__expand_as_type(scope));
scope.register(Instruction::new(Metadata::Shape { dim, list }, out));
out.into()
})
}
pub fn coordinate(&self, index: usize, dim: usize) -> usize {
intrinsic!(|scope| {
let index: Value = index.into();
let stride = self.__expand_stride_method(scope, dim.clone());
let shape = self.__expand_shape_method(scope, dim.clone());
let num_strides = scope.create_value(usize::__expand_as_type(scope));
scope.register(Instruction::new(
Arithmetic::Div(BinaryOperands {
lhs: index,
rhs: stride.expand.into(),
}),
num_strides.clone().into(),
));
let coordinate = scope.create_value(usize::__expand_as_type(scope));
scope.register(Instruction::new(
Arithmetic::Rem(BinaryOperands {
lhs: num_strides,
rhs: shape.expand.into(),
}),
coordinate.clone().into(),
));
coordinate.into()
})
}
#[allow(clippy::len_without_is_empty)]
pub fn len(&self) -> usize {
self.meta.len
}
#[allow(clippy::len_without_is_empty)]
pub fn buffer_len(&self) -> usize {
intrinsic!(|scope| { self.__extract_length(scope) })
}
pub fn rank(&self) -> usize {
self.meta.rank
}
}
}
mod vector {
use super::*;
impl<P: Scalar, N: Size> Tensor<Vector<P, N>> {
pub fn vector_size(&self) -> VectorSize {
N::value()
}
pub fn __expand_vector_size(
expand: <Self as CubeType>::ExpandType,
scope: &Scope,
) -> VectorSize {
expand.__expand_vector_size_method(scope)
}
}
}
impl<'a, E: CubePrimitive> From<&'a OwnedTensorExpand<E>> for &'a TensorExpand<E> {
fn from(value: &'a OwnedTensorExpand<E>) -> Self {
value.deref()
}
}
impl<'a, E: CubePrimitive> From<&'a mut OwnedTensorExpand<E>> for &'a mut TensorExpand<E> {
fn from(value: &'a mut OwnedTensorExpand<E>) -> Self {
value.deref_mut()
}
}
impl<'a, E: CubePrimitive> From<&'a TensorExpand<E>> for &'a SliceExpand<E> {
fn from(value: &'a TensorExpand<E>) -> Self {
value.deref()
}
}
impl<'a, E: CubePrimitive> From<&'a mut TensorExpand<E>> for &'a mut SliceExpand<E> {
fn from(value: &'a mut TensorExpand<E>) -> Self {
value.deref_mut()
}
}
impl<T: CubePrimitive> SizedContainer<usize> for Tensor<T> {
fn len(&self) -> usize {
unexpanded!()
}
}
impl<T: CubePrimitive> SizedContainerExpand<usize> for TensorExpand<T> {
fn __expand_len_method(&self, scope: &Scope) -> NativeExpand<usize> {
self.__expand_len_method(scope)
}
}
impl<T: CubePrimitive> Iterator for &Tensor<T> {
type Item = T;
fn next(&mut self) -> Option<Self::Item> {
unexpanded!()
}
}
impl<T: CubePrimitive> List<T> for Tensor<T> {}
impl<T: CubePrimitive> Deref for Tensor<T> {
type Target = [T];
fn deref(&self) -> &Self::Target {
unexpanded!()
}
}
impl<T: CubePrimitive> DerefMut for Tensor<T> {
fn deref_mut(&mut self) -> &mut Self::Target {
unexpanded!()
}
}
impl<T: CubePrimitive> Deref for TensorExpand<T> {
type Target = SliceExpand<T>;
fn deref(&self) -> &Self::Target {
&self.buffer
}
}
impl<T: CubePrimitive> DerefMut for TensorExpand<T> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.buffer
}
}
impl<T: CubePrimitive> Deref for OwnedTensor<T> {
type Target = Tensor<T>;
fn deref(&self) -> &Self::Target {
unexpanded!()
}
}
impl<T: CubePrimitive> DerefMut for OwnedTensor<T> {
fn deref_mut(&mut self) -> &mut Self::Target {
unexpanded!()
}
}
impl<T: CubePrimitive> Deref for OwnedTensorExpand<T> {
type Target = TensorExpand<T>;
fn deref(&self) -> &Self::Target {
unsafe { core::mem::transmute(self) }
}
}
impl<T: CubePrimitive> DerefMut for OwnedTensorExpand<T> {
fn deref_mut(&mut self) -> &mut Self::Target {
unsafe { core::mem::transmute(self) }
}
}
impl<T: CubePrimitive> ListExpand<T> for TensorExpand<T> {
fn __expand_len_method(&self, scope: &Scope) -> NativeExpand<usize> {
Self::__expand_len_method(self, scope)
}
}
impl<T: CubePrimitive> Vectorized for Tensor<T> {}
impl<T: CubePrimitive> VectorizedExpand for TensorExpand<T> {
fn vector_size(&self) -> VectorSize {
self.buffer.expand.ty.vector_size()
}
}