use crate::{Result, ShapeVec, StrideVec, ValidationError};
use std::fmt::Debug;
pub trait TensorRank: private::Sealed + Clone + Copy + Debug + Eq + Send + Sync + 'static {
const RANK: Option<usize>;
type Shape: Clone + Debug + PartialEq + Eq + AsRef<[usize]>;
type Strides: Clone + Debug + PartialEq + Eq + AsRef<[isize]>;
fn shape_from_vec(shape: ShapeVec) -> Result<Self::Shape>;
fn shape_into_vec(shape: Self::Shape) -> ShapeVec;
fn strides_from_vec(strides: StrideVec) -> Result<Self::Strides>;
fn strides_into_vec(strides: Self::Strides) -> StrideVec;
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct DynRank;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub struct Rank<const N: usize>;
impl TensorRank for DynRank {
const RANK: Option<usize> = None;
type Shape = ShapeVec;
type Strides = StrideVec;
fn shape_from_vec(shape: ShapeVec) -> Result<Self::Shape> {
Ok(shape)
}
fn shape_into_vec(shape: Self::Shape) -> ShapeVec {
shape
}
fn strides_from_vec(strides: StrideVec) -> Result<Self::Strides> {
Ok(strides)
}
fn strides_into_vec(strides: Self::Strides) -> StrideVec {
strides
}
}
impl<const N: usize> TensorRank for Rank<N> {
const RANK: Option<usize> = Some(N);
type Shape = [usize; N];
type Strides = [isize; N];
fn shape_from_vec(shape: ShapeVec) -> Result<Self::Shape> {
let actual = shape.len();
shape
.into_vec()
.try_into()
.map_err(|_| ValidationError::RankMismatch {
expected: N,
actual,
})
}
fn shape_into_vec(shape: Self::Shape) -> ShapeVec {
ShapeVec::from_iter(shape)
}
fn strides_from_vec(strides: StrideVec) -> Result<Self::Strides> {
let actual = strides.len();
strides
.into_vec()
.try_into()
.map_err(|_| ValidationError::RankMismatch {
expected: N,
actual,
})
}
fn strides_into_vec(strides: Self::Strides) -> StrideVec {
StrideVec::from_iter(strides)
}
}
mod private {
pub trait Sealed {}
impl Sealed for super::DynRank {}
impl<const N: usize> Sealed for super::Rank<N> {}
}