tensor4all_tensorbackend/
tensor_element.rs1use anyhow::{anyhow, ensure, Result};
2use num_complex::{Complex32, Complex64};
3use tenferro::{DType, Tensor as NativeTensor, TensorScalar};
4
5pub trait TensorElement: TensorScalar + Copy + Send + Sync + 'static {
21 fn dense_native_tensor_from_col_major(data: &[Self], dims: &[usize]) -> Result<NativeTensor>;
28
29 fn diag_native_tensor_from_col_major(
36 data: &[Self],
37 logical_rank: usize,
38 ) -> Result<NativeTensor>;
39
40 fn scalar_native_tensor(value: Self) -> Result<NativeTensor>;
47
48 fn dense_values_from_native_col_major(tensor: &NativeTensor) -> Result<Vec<Self>>;
55
56 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);