tenferro_internal_cpu_kernels/
lib.rs1#![doc(hidden)]
2
3#[allow(unused_macros)]
10macro_rules! preset_scalar {
11 (F32) => {
12 f32
13 };
14 (F64) => {
15 f64
16 };
17 (I32) => {
18 i32
19 };
20 (I64) => {
21 i64
22 };
23 (Bool) => {
24 bool
25 };
26 (C32) => {
27 num_complex::Complex32
28 };
29 (C64) => {
30 num_complex::Complex64
31 };
32}
33pub type Result<T> = tenferro_tensor::Result<T>;
34pub use tenferro_cpu_basic::{
35 cpu_backend_buffer_error, cpu_division_by_zero, typed_host_data, typed_view,
36 typed_view_from_view, ConjElem, CpuNumericalError,
37};
38pub use tenferro_tensor::{CacheStats, DType, Error, ErrorKind};
39
40pub mod dispatch;
41pub mod elementwise;
42pub mod read_into;
43pub mod scalar_ops;
44pub use read_into::elementwise_read_into_with_context;
45
46pub use num_complex::{Complex32, Complex64};
48#[cfg(test)]
49use std::mem::MaybeUninit;
50#[cfg(test)]
51use strided_kernel::{map_into, Identity, StridedView};
52#[cfg(test)]
53use tenferro_cpu_basic::{BufferPool, PoolScalar, PooledUninitOutput};
54#[cfg(test)]
55use tenferro_tensor::{Tensor, TensorRank, TensorRead, TensorView, TypedTensor, TypedTensorView};
56
57#[cfg(test)]
58pub(crate) fn materialize_tensor_read(
59 buffers: &mut BufferPool,
60 op: &'static str,
61 input: TensorRead<'_>,
62) -> Result<Tensor> {
63 match input {
64 TensorRead::Tensor(tensor) => clone_host_tensor_read(op, tensor),
65 TensorRead::View(view) => materialize_tensor_view(buffers, op, view),
66 }
67}
68
69#[cfg(test)]
70fn clone_host_tensor_read(op: &'static str, tensor: &Tensor) -> Result<Tensor> {
71 macro_rules! clone_host {
72 ($variant:ident, $tensor:expr) => {{
73 typed_host_data(op, $tensor)?;
74 Ok(Tensor::from_typed::<preset_scalar!($variant)>(
75 $tensor.duplicate()?,
76 ))
77 }};
78 }
79 match tensor.dtype() {
80 DType::F32 => {
81 let tensor = host_typed::<f32>(op, tensor)?;
82 clone_host!(F32, tensor)
83 }
84 DType::F64 => {
85 let tensor = host_typed::<f64>(op, tensor)?;
86 clone_host!(F64, tensor)
87 }
88 DType::I32 => {
89 let tensor = host_typed::<i32>(op, tensor)?;
90 clone_host!(I32, tensor)
91 }
92 DType::I64 => {
93 let tensor = host_typed::<i64>(op, tensor)?;
94 clone_host!(I64, tensor)
95 }
96 DType::Bool => {
97 let tensor = host_typed::<bool>(op, tensor)?;
98 clone_host!(Bool, tensor)
99 }
100 DType::C32 => {
101 let tensor = host_typed::<Complex32>(op, tensor)?;
102 clone_host!(C32, tensor)
103 }
104 DType::C64 => {
105 let tensor = host_typed::<Complex64>(op, tensor)?;
106 clone_host!(C64, tensor)
107 }
108 DType::External(type_id) => Err(crate::Error::unsupported_dtype(
111 op,
112 tenferro_tensor::DType::External(type_id),
113 "an externally defined payload must be duplicated by its owner",
114 )),
115 }
116}
117
118#[cfg(test)]
119fn host_typed<'a, T: tenferro_tensor::TensorScalar>(
125 op: &'static str,
126 tensor: &'a Tensor,
127) -> Result<&'a TypedTensor<T>> {
128 tensor.as_typed::<T>().ok_or_else(|| {
129 crate::Error::unsupported_dtype(
130 op,
131 tensor.dtype(),
132 "an externally defined payload must be duplicated by its owner",
133 )
134 })
135}
136
137#[cfg(test)]
138fn materialize_tensor_view(
139 buffers: &mut BufferPool,
140 op: &'static str,
141 view: TensorView<'_>,
142) -> Result<Tensor> {
143 macro_rules! materialize {
144 ($variant:ident, $view:expr) => {{
145 Ok(Tensor::from_typed::<preset_scalar!($variant)>(
146 typed_materialize_view_for_tests(buffers, &$view, op)?,
147 ))
148 }};
149 }
150 match view {
151 TensorView::F32(view) => materialize!(F32, view),
152 TensorView::F64(view) => materialize!(F64, view),
153 TensorView::I32(view) => materialize!(I32, view),
154 TensorView::I64(view) => materialize!(I64, view),
155 TensorView::Bool(view) => materialize!(Bool, view),
156 TensorView::C32(view) => materialize!(C32, view),
157 TensorView::C64(view) => materialize!(C64, view),
158 }
159}
160
161#[cfg(test)]
162fn typed_materialize_view_for_tests<T, R>(
163 buffers: &mut BufferPool,
164 view: &TypedTensorView<'_, T, R>,
165 op: &'static str,
166) -> Result<TypedTensor<T, R>>
167where
168 T: Copy + Clone + PoolScalar + 'static,
169 R: TensorRank,
170{
171 if view.backend_buffer().is_some() {
172 return Err(cpu_backend_buffer_error(op));
173 }
174 let src: StridedView<'_, T, Identity> = StridedView::new(
175 view.host_storage()?,
176 view.shape(),
177 view.strides(),
178 view.offset(),
179 )
180 .map_err(|err| Error::backend_source(op, err))?;
181 let mut out = PooledUninitOutput::<T>::new(buffers, view.shape().to_vec())?;
182 map_into(&mut out.as_uninit_view_mut()?, &src, |x| {
183 MaybeUninit::new(x)
184 })
185 .map_err(|err| Error::backend_source(op, err))?;
186 let out = unsafe { out.assume_init_as::<R>()? };
188 let shape = R::shape_from_vec(view.shape().to_vec().into())
189 .map_err(|err| Error::backend_source(op, err))?;
190 let mut tensor = TypedTensor::from_vec_col_major(
191 shape,
192 out.into_vec_col_major()
193 .map_err(|failure| failure.into_parts().1)?
194 .1,
195 )?;
196 tensor.set_placement(view.placement().clone());
197 Ok(tensor)
198}
199
200#[cfg(test)]
201mod tests;