1use 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
38pub type ShapeVec = SmallVec<[usize; 8]>;
49
50pub trait IntoShapeVec {
55 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
68pub type StrideVec = SmallVec<[isize; 8]>;
79
80pub type Result<T> = std::result::Result<T, ValidationError>;
91
92define_scalar_tag! {
93 pub enum DType {
103 F32 => f32 : Float 0 32,
105 F64 => f64 : Float 1 64,
107 I32 => i32 : Integer 0 32,
109 I64 => i64 : Integer 1 64,
111 Bool => bool : Boolean 0 0,
113 C32 => Complex32 : Complex 0 32,
115 C64 => Complex64 : Complex 1 64,
117 }
118 external External(core::any::TypeId);
119}
120
121pub trait TensorScalar: Scalar + Copy + Clone + Send + Sync + 'static + private::Sealed {
132 type Real: TensorScalar;
134
135 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
265impl_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#[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
310pub 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 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}