use core::ops::{Index, IndexMut};
use crate::Shape;
use crate::element::Element;
use crate::indexing::AsIndex;
use crate::tensor::{DType, ravel_index};
use super::{DataError, TensorData};
impl TensorData {
pub fn try_view<E: Element>(&self) -> Result<TensorDataView<'_, E>, DataError> {
TensorDataView::<E>::try_view(self)
}
#[track_caller]
pub fn view<E: Element>(&self) -> TensorDataView<'_, E> {
self.try_view()
.unwrap_or_else(|err| panic!("Failed to create TensorData view: {err}"))
}
pub fn try_mut_view<E: Element>(&mut self) -> Result<TensorDataViewMut<'_, E>, DataError> {
TensorDataViewMut::<E>::try_mut_view(self)
}
#[track_caller]
pub fn mut_view<E: Element>(&mut self) -> TensorDataViewMut<'_, E> {
self.try_mut_view()
.unwrap_or_else(|err| panic!("Failed to create mutable TensorData view: {err}"))
}
}
#[derive(Debug)]
pub struct TensorDataView<'a, E: Element> {
values: &'a [E],
shape: &'a Shape,
dtype: DType,
}
impl<'a, E: Element> TensorDataView<'a, E> {
pub fn try_view(data: &'a TensorData) -> Result<TensorDataView<'a, E>, DataError> {
let shape = &data.shape;
let dtype = data.dtype;
let expected = shape.num_elements();
let values = data.as_slice::<E>()?;
let actual = values.len();
if actual != expected {
return Err(DataError::ElementCountMismatch { expected, actual });
}
Ok(TensorDataView {
values,
shape,
dtype,
})
}
pub fn shape(&self) -> &Shape {
self.shape
}
pub fn dtype(&self) -> DType {
self.dtype
}
pub fn ravel_index<I: AsIndex>(&self, index: &[I]) -> usize {
ravel_index(index, self.shape)
}
}
impl<'a, I: AsIndex, E: Element> Index<&[I]> for TensorDataView<'a, E> {
type Output = E;
fn index(&self, index: &[I]) -> &Self::Output {
let o = self.ravel_index(index);
&self.values[o]
}
}
#[derive(Debug)]
pub struct TensorDataViewMut<'a, E: Element> {
values: &'a mut [E],
shape: Shape,
dtype: DType,
}
impl<'a, E: Element> TensorDataViewMut<'a, E> {
pub fn try_mut_view(data: &'a mut TensorData) -> Result<TensorDataViewMut<'a, E>, DataError> {
let shape = data.shape.clone();
let dtype = data.dtype;
let expected = shape.num_elements();
let values = data.as_mut_slice::<E>()?;
let actual = values.len();
if actual != expected {
return Err(DataError::ElementCountMismatch { expected, actual });
}
Ok(TensorDataViewMut {
values,
shape,
dtype,
})
}
pub fn shape(&self) -> &Shape {
&self.shape
}
pub fn dtype(&self) -> DType {
self.dtype
}
pub fn ravel_index<I: AsIndex>(&self, index: &[I]) -> usize {
ravel_index(index, &self.shape)
}
}
impl<'a, I, E> Index<&[I]> for TensorDataViewMut<'a, E>
where
I: AsIndex,
E: Element,
{
type Output = E;
fn index(&self, index: &[I]) -> &Self::Output {
let o = self.ravel_index::<I>(index);
&self.values[o]
}
}
impl<'a, I, E> IndexMut<&[I]> for TensorDataViewMut<'a, E>
where
I: AsIndex,
E: Element,
{
fn index_mut(&mut self, index: &[I]) -> &mut Self::Output {
let o = self.ravel_index::<I>(index);
&mut self.values[o]
}
}