Skip to main content

tenferro_tensor_core/
lib.rs

1//! Backend-independent rank/layout, dtype and scalar metadata.
2//!
3//! `tenferro-tensor-core` owns backend-independent tensor metadata: scalar
4//! tags and promotion facts, rank/layout metadata, and layout validation. It
5//! does not own tensor storage, execution backends, backend buffers, GPU
6//! handles, provider selection, or materializing kernels. The tensor families
7//! that carry storage live in `tenferro-tensor`.
8//!
9//! # Examples
10//!
11//! ```rust
12//! use tenferro_tensor_core::{DType, Rank, TensorLayout};
13//!
14//! let layout = TensorLayout::<Rank<2>>::compact([2, 3])?;
15//! let transposed = layout.transpose_view([1, 0])?;
16//! assert_eq!(transposed.shape(), &[3, 2]);
17//!
18//! assert_eq!(DType::F64.spec().width, 64);
19//! # Ok::<(), tenferro_tensor_core::ValidationError>(())
20//! ```
21
22use num_complex::{Complex32, Complex64};
23use smallvec::SmallVec;
24
25mod error;
26mod layout;
27mod rank;
28mod scalar;
29#[macro_use]
30mod scalar_set;
31
32pub use error::{ErrorKind, ShapeMismatch, ValidationError, ValidationKind};
33pub use layout::TensorLayout;
34pub use rank::{DynRank, IntoRankShape, Rank, TensorRank};
35pub use scalar::{ad_admission, AdAdmissionError, Scalar, ScalarArithmetic, ScalarDomain};
36pub use scalar_set::{promote_in_set, promote_specs, MemberKind, MemberSpec};
37
38/// Small tensor shape vector with inline capacity for common dynamic ranks.
39///
40/// # Examples
41///
42/// ```rust
43/// use tenferro_tensor_core::ShapeVec;
44///
45/// let shape = ShapeVec::from_vec(vec![2, 3]);
46/// assert_eq!(shape.as_slice(), &[2, 3]);
47/// ```
48pub type ShapeVec = SmallVec<[usize; 8]>;
49
50/// Convert a common shape container into the dynamic owned shape type.
51///
52/// Arrays, vectors, slices, and [`ShapeVec`] implement this trait through
53/// their `AsRef<[usize]>` representation.
54pub trait IntoShapeVec {
55    /// Convert this shape container into an owned [`ShapeVec`].
56    fn into_shape_vec(self) -> ShapeVec;
57}
58
59impl<S> IntoShapeVec for S
60where
61    S: AsRef<[usize]>,
62{
63    fn into_shape_vec(self) -> ShapeVec {
64        self.as_ref().iter().copied().collect()
65    }
66}
67
68/// Small tensor stride vector with signed element strides.
69///
70/// # Examples
71///
72/// ```rust
73/// use tenferro_tensor_core::StrideVec;
74///
75/// let strides = StrideVec::from_vec(vec![1, 2]);
76/// assert_eq!(strides.as_slice(), &[1, 2]);
77/// ```
78pub type StrideVec = SmallVec<[isize; 8]>;
79
80/// Result type for tensor data-model operations.
81///
82/// # Examples
83///
84/// ```rust
85/// use tenferro_tensor_core::{Result, ValidationError};
86///
87/// let result: Result<()> = Err(ValidationError::RankMismatch { expected: 2, actual: 1 });
88/// assert!(result.is_err());
89/// ```
90pub type Result<T> = std::result::Result<T, ValidationError>;
91
92define_scalar_tag! {
93    /// Runtime scalar dtype tag.
94    ///
95    /// # Examples
96    ///
97    /// ```rust
98    /// use tenferro_tensor_core::DType;
99    ///
100    /// assert_eq!(DType::F64, DType::F64);
101    /// ```
102    pub enum DType {
103        /// 32-bit floating point.
104        F32 => f32 : Float 0 32,
105        /// 64-bit floating point.
106        F64 => f64 : Float 1 64,
107        /// 32-bit signed integer.
108        I32 => i32 : Integer 0 32,
109        /// 64-bit signed integer.
110        I64 => i64 : Integer 1 64,
111        /// Boolean.
112        Bool => bool : Boolean 0 0,
113        /// 32-bit complex floating point.
114        C32 => Complex32 : Complex 0 32,
115        /// 64-bit complex floating point.
116        C64 => Complex64 : Complex 1 64,
117    }
118    external External(core::any::TypeId);
119}
120
121/// Sealed trait for scalar types supported by the core tensor data model.
122///
123/// # Examples
124///
125/// ```rust
126/// use tenferro_tensor_core::{DType, TensorScalar};
127///
128/// assert_eq!(f64::dtype(), DType::F64);
129/// assert_eq!(num_complex::Complex64::dtype(), DType::C64);
130/// ```
131pub trait TensorScalar: Scalar + Copy + Clone + Send + Sync + 'static + private::Sealed {
132    /// Real-valued counterpart of this scalar type.
133    type Real: TensorScalar;
134
135    /// Return the scalar dtype tag.
136    ///
137    /// # Examples
138    ///
139    /// ```rust
140    /// use tenferro_tensor_core::{DType, TensorScalar};
141    ///
142    /// assert_eq!(i64::dtype(), DType::I64);
143    /// ```
144    fn dtype() -> DType;
145}
146
147mod private {
148    pub trait Sealed {}
149
150    impl Sealed for f32 {}
151    impl Sealed for f64 {}
152    impl Sealed for i32 {}
153    impl Sealed for i64 {}
154    impl Sealed for bool {}
155    impl Sealed for num_complex::Complex32 {}
156    impl Sealed for num_complex::Complex64 {}
157}
158
159macro_rules! scalar_domain {
160    (float) => {
161        ScalarDomain::Field
162    };
163    (complex) => {
164        ScalarDomain::Field
165    };
166    (integer) => {
167        ScalarDomain::Field
168    };
169    (boolean) => {
170        ScalarDomain::NonField
171    };
172}
173
174macro_rules! impl_scalar_arithmetic {
175    ($ty:ty, float) => {
176        impl ScalarArithmetic for $ty {
177            fn scalar_zero() -> Self {
178                0.0
179            }
180
181            fn scalar_one() -> Self {
182                1.0
183            }
184
185            fn scalar_add(self, rhs: Self) -> Self {
186                std::ops::Add::add(self, rhs)
187            }
188
189            fn scalar_sub(self, rhs: Self) -> Self {
190                std::ops::Sub::sub(self, rhs)
191            }
192
193            fn scalar_mul(self, rhs: Self) -> Self {
194                std::ops::Mul::mul(self, rhs)
195            }
196        }
197    };
198    ($ty:ty, complex) => {
199        impl ScalarArithmetic for $ty {
200            fn scalar_zero() -> Self {
201                <$ty>::new(0.0, 0.0)
202            }
203
204            fn scalar_one() -> Self {
205                <$ty>::new(1.0, 0.0)
206            }
207
208            fn scalar_add(self, rhs: Self) -> Self {
209                std::ops::Add::add(self, rhs)
210            }
211
212            fn scalar_sub(self, rhs: Self) -> Self {
213                std::ops::Sub::sub(self, rhs)
214            }
215
216            fn scalar_mul(self, rhs: Self) -> Self {
217                std::ops::Mul::mul(self, rhs)
218            }
219        }
220    };
221    ($ty:ty, integer) => {
222        impl ScalarArithmetic for $ty {
223            fn scalar_zero() -> Self {
224                0
225            }
226
227            fn scalar_one() -> Self {
228                1
229            }
230
231            fn scalar_add(self, rhs: Self) -> Self {
232                self.wrapping_add(rhs)
233            }
234
235            fn scalar_sub(self, rhs: Self) -> Self {
236                self.wrapping_sub(rhs)
237            }
238
239            fn scalar_mul(self, rhs: Self) -> Self {
240                self.wrapping_mul(rhs)
241            }
242        }
243    };
244    ($ty:ty, boolean) => {};
245}
246
247macro_rules! impl_scalar {
248    ($ty:ty, $real:ty, $dtype:expr, $variant:ident, $kind:ident) => {
249        impl Scalar for $ty {
250            const DOMAIN: ScalarDomain = scalar_domain!($kind);
251        }
252
253        impl_scalar_arithmetic!($ty, $kind);
254
255        impl TensorScalar for $ty {
256            type Real = $real;
257
258            fn dtype() -> DType {
259                $dtype
260            }
261        }
262    };
263}
264
265// The single preset table: tag, real counterpart, erased variant, and algebra
266// kind are declared once here and expanded into every contract the preset
267// scalars implement.
268impl_scalar!(f32, f32, DType::F32, F32, float);
269impl_scalar!(f64, f64, DType::F64, F64, float);
270impl_scalar!(i32, i32, DType::I32, I32, integer);
271impl_scalar!(i64, i64, DType::I64, I64, integer);
272impl_scalar!(bool, bool, DType::Bool, Bool, boolean);
273impl_scalar!(Complex32, f32, DType::C32, C32, complex);
274impl_scalar!(Complex64, f64, DType::C64, C64, complex);
275
276/// Explicit slice descriptor.
277///
278/// A zero step is invalid. Layout metadata APIs support signed steps when
279/// reachable-range validation proves the view stays inside the backing
280/// allocation.
281///
282/// # Examples
283///
284/// ```rust
285/// use tenferro_tensor_core::SliceSpec;
286///
287/// let spec = SliceSpec { start: 1, end: 4, step: 2 };
288/// assert_eq!(spec.step, 2);
289/// ```
290#[derive(Clone, Copy, Debug, PartialEq, Eq)]
291pub struct SliceSpec {
292    pub start: isize,
293    pub end: isize,
294    pub step: isize,
295}
296
297fn checked_product(shape: &[usize]) -> Result<usize> {
298    shape.iter().try_fold(1usize, |acc, &dim| {
299        acc.checked_mul(dim).ok_or(ValidationError::IntegerOverflow)
300    })
301}
302
303fn checked_logical_element_count(shape: &[usize]) -> Result<usize> {
304    if shape.contains(&0) {
305        return Ok(0);
306    }
307    checked_product(shape)
308}
309
310/// Return compact column-major strides for a shape.
311///
312/// # Examples
313///
314/// ```rust
315/// use tenferro_tensor_core::col_major_strides;
316///
317/// assert_eq!(col_major_strides(&[2, 3])?.as_slice(), &[1, 2]);
318/// # Ok::<(), tenferro_tensor_core::ValidationError>(())
319/// ```
320///
321/// # Errors
322///
323/// Returns [`ValidationError::IntegerOverflow`] when a stride or extent
324/// cannot be represented by the metadata arithmetic.
325pub fn col_major_strides(shape: &[usize]) -> Result<StrideVec> {
326    let mut strides = StrideVec::new();
327    let mut stride = 1isize;
328    for &extent in shape {
329        strides.push(stride);
330        let extent = isize::try_from(extent).map_err(|_| ValidationError::IntegerOverflow)?;
331        stride = stride
332            .checked_mul(extent)
333            .ok_or(ValidationError::IntegerOverflow)?;
334    }
335    Ok(strides)
336}
337
338fn validate_permutation(rank: usize, axes: &[usize]) -> Result<()> {
339    if axes.len() != rank {
340        return Err(ValidationError::InvalidPermutationLength {
341            expected: rank,
342            actual: axes.len(),
343        });
344    }
345    // Common ranks track seen axes in a bitmask; only very high ranks allocate.
346    if rank <= u128::BITS as usize {
347        let mut seen = 0u128;
348        for &axis in axes {
349            if axis >= rank {
350                return Err(ValidationError::AxisOutOfBounds { axis, rank });
351            }
352            let bit = 1u128 << axis;
353            if seen & bit != 0 {
354                return Err(ValidationError::DuplicateAxis {
355                    axis,
356                    role: "permutation",
357                });
358            }
359            seen |= bit;
360        }
361        return Ok(());
362    }
363    let mut seen = vec![false; rank];
364    for &axis in axes {
365        if axis >= rank {
366            return Err(ValidationError::AxisOutOfBounds { axis, rank });
367        }
368        if seen[axis] {
369            return Err(ValidationError::DuplicateAxis {
370                axis,
371                role: "permutation",
372            });
373        }
374        seen[axis] = true;
375    }
376    Ok(())
377}