Skip to main content

tensor4all_tensorbackend/
tensor_element.rs

1use anyhow::{anyhow, ensure, Result};
2use num_complex::{Complex32, Complex64};
3use tenferro::{DType, Tensor as NativeTensor, TensorScalar};
4
5/// Public scalar element types supported by tensor4all dense/diag constructors.
6///
7/// Implemented for `f32`, `f64`, `Complex32`, and `Complex64`.
8///
9/// # Examples
10///
11/// ```
12/// use tensor4all_tensorbackend::TensorElement;
13///
14/// let t = f64::dense_native_tensor_from_col_major(&[1.0, 2.0], &[2]).unwrap();
15/// assert_eq!(t.shape(), &[2]);
16///
17/// let vals = f64::dense_values_from_native_col_major(&t).unwrap();
18/// assert_eq!(vals, vec![1.0, 2.0]);
19/// ```
20pub trait TensorElement: TensorScalar + Copy + Send + Sync + 'static {
21    /// Build a dense native tensor from column-major data.
22    /// # Errors
23    ///
24    /// Returns an error when the data length does not match the dimensions (a
25    /// /// shape mismatch) or the backend conversion fails.
26    ///
27    fn dense_native_tensor_from_col_major(data: &[Self], dims: &[usize]) -> Result<NativeTensor>;
28
29    /// Build a diagonal native tensor from column-major diagonal payload data.
30    /// # Errors
31    ///
32    /// Returns an error when the diagonal payload is incompatible (a shape
33    /// /// mismatch) or the backend conversion fails.
34    ///
35    fn diag_native_tensor_from_col_major(
36        data: &[Self],
37        logical_rank: usize,
38    ) -> Result<NativeTensor>;
39
40    /// Build a rank-0 native tensor.
41    /// # Errors
42    ///
43    /// Returns an error when the scalar cannot be converted (a dtype mismatch or
44    /// /// backend failure).
45    ///
46    fn scalar_native_tensor(value: Self) -> Result<NativeTensor>;
47
48    /// Materialize dense column-major values from a native tensor.
49    /// # Errors
50    ///
51    /// Returns an error when the native tensor cannot be materialized (a dtype
52    /// /// mismatch or backend failure).
53    ///
54    fn dense_values_from_native_col_major(tensor: &NativeTensor) -> Result<Vec<Self>>;
55
56    /// Materialize diagonal values from a dense native tensor.
57    /// # Errors
58    ///
59    /// Returns an error when the native tensor cannot be materialized (a dtype
60    /// /// mismatch or backend failure).
61    ///
62    fn diag_values_from_native_temp(tensor: &NativeTensor) -> Result<Vec<Self>>;
63}
64
65fn tensor_dtype_name(dtype: DType) -> &'static str {
66    match dtype {
67        DType::F32 => "f32",
68        DType::F64 => "f64",
69        DType::I32 => "i32",
70        DType::I64 => "i64",
71        DType::Bool => "bool",
72        DType::C32 => "c32",
73        DType::C64 => "c64",
74    }
75}
76
77fn checked_product(dims: &[usize]) -> Result<usize> {
78    dims.iter().try_fold(1usize, |acc, &dim| {
79        acc.checked_mul(dim)
80            .ok_or_else(|| anyhow::anyhow!("dimension product overflow"))
81    })
82}
83
84fn dense_diagonal_values<T: Copy + Default>(diag: &[T], logical_rank: usize) -> Result<Vec<T>> {
85    ensure!(
86        logical_rank >= 1,
87        "diagonal tensor construction requires at least one logical axis"
88    );
89    let diag_len = diag.len();
90    let dims = vec![diag_len; logical_rank];
91    let total_len = checked_product(&dims)?;
92    let mut dense = vec![T::default(); total_len];
93    let diagonal_stride = (0..logical_rank)
94        .scan(1usize, |stride, _| {
95            let current = *stride;
96            *stride = stride.saturating_mul(diag_len);
97            Some(current)
98        })
99        .sum::<usize>();
100    for (i, value) in diag.iter().copied().enumerate() {
101        dense[i * diagonal_stride] = value;
102    }
103    Ok(dense)
104}
105
106macro_rules! impl_tensor_element {
107    ($ty:ty, $dtype:expr) => {
108        impl TensorElement for $ty {
109            fn dense_native_tensor_from_col_major(
110                data: &[Self],
111                dims: &[usize],
112            ) -> Result<NativeTensor> {
113                let expected_len: usize = checked_product(dims)?;
114                ensure!(
115                    data.len() == expected_len,
116                    "dense tensor len {} does not match dims {:?} (expected {})",
117                    data.len(),
118                    dims,
119                    expected_len
120                );
121                Ok(NativeTensor::from_vec_col_major(
122                    dims.to_vec(),
123                    data.to_vec(),
124                )?)
125            }
126
127            fn diag_native_tensor_from_col_major(
128                data: &[Self],
129                logical_rank: usize,
130            ) -> Result<NativeTensor> {
131                let dims = vec![data.len(); logical_rank];
132                let dense = dense_diagonal_values(data, logical_rank)?;
133                Self::dense_native_tensor_from_col_major(&dense, &dims)
134            }
135
136            fn scalar_native_tensor(value: Self) -> Result<NativeTensor> {
137                Ok(NativeTensor::from_vec_col_major(vec![], vec![value])?)
138            }
139
140            fn dense_values_from_native_col_major(tensor: &NativeTensor) -> Result<Vec<Self>> {
141                tensor
142                    .as_slice::<Self>()
143                    .map(|values| values.to_vec())
144                    .map_err(|_| {
145                        anyhow!(
146                            "tensor dtype mismatch: expected {}, got {}",
147                            tensor_dtype_name($dtype),
148                            tensor_dtype_name(tensor.dtype())
149                        )
150                    })
151            }
152
153            fn diag_values_from_native_temp(tensor: &NativeTensor) -> Result<Vec<Self>> {
154                let shape = tensor.shape();
155                ensure!(
156                    !shape.is_empty(),
157                    "diagonal extraction requires rank >= 1, got scalar tensor"
158                );
159                let diag_len = shape[0];
160                ensure!(
161                    shape.iter().all(|&dim| dim == diag_len),
162                    "expected square/equal dims for diagonal extraction, got {:?}",
163                    shape
164                );
165                let dense = Self::dense_values_from_native_col_major(tensor)?;
166                let diagonal_stride = (0..shape.len())
167                    .scan(1usize, |stride, _| {
168                        let current = *stride;
169                        *stride = stride.saturating_mul(diag_len);
170                        Some(current)
171                    })
172                    .sum::<usize>();
173                Ok((0..diag_len).map(|i| dense[i * diagonal_stride]).collect())
174            }
175        }
176    };
177}
178
179impl_tensor_element!(f32, DType::F32);
180impl_tensor_element!(f64, DType::F64);
181impl_tensor_element!(Complex32, DType::C32);
182impl_tensor_element!(Complex64, DType::C64);