tenferro_internal_cpu_kernels/
lib.rs1#![doc(hidden)]
2
3pub type Result<T> = tenferro_tensor::Result<T>;
9pub use tenferro_cpu_basic::{
10 cpu_backend_buffer_error, cpu_division_by_zero, typed_host_data, typed_view,
11 typed_view_from_view, ConjElem, CpuNumericalError,
12};
13pub use tenferro_tensor::{CacheStats, DType, Error, ErrorKind};
14
15pub mod elementwise;
16pub mod read_into;
17pub use read_into::elementwise_read_into_with_context;
18
19#[cfg(test)]
20use std::mem::MaybeUninit;
21#[cfg(test)]
22use strided_kernel::{map_into, Identity, StridedView};
23#[cfg(test)]
24use tenferro_cpu_basic::{BufferPool, PoolScalar, PooledUninitOutput};
25#[cfg(test)]
26use tenferro_tensor::{Tensor, TensorRank, TensorRead, TensorView, TypedTensor, TypedTensorView};
27
28#[cfg(test)]
29pub(crate) fn materialize_tensor_read(
30 buffers: &mut BufferPool,
31 op: &'static str,
32 input: TensorRead<'_>,
33) -> Result<Tensor> {
34 match input {
35 TensorRead::Tensor(tensor) => clone_host_tensor_read(op, tensor),
36 TensorRead::View(view) => materialize_tensor_view(buffers, op, view),
37 }
38}
39
40#[cfg(test)]
41fn clone_host_tensor_read(op: &'static str, tensor: &Tensor) -> Result<Tensor> {
42 macro_rules! clone_host {
43 ($variant:ident, $tensor:expr) => {{
44 typed_host_data(op, $tensor)?;
45 Ok(Tensor::$variant($tensor.duplicate()?))
46 }};
47 }
48 match tensor {
49 Tensor::F32(tensor) => clone_host!(F32, tensor),
50 Tensor::F64(tensor) => clone_host!(F64, tensor),
51 Tensor::I32(tensor) => clone_host!(I32, tensor),
52 Tensor::I64(tensor) => clone_host!(I64, tensor),
53 Tensor::Bool(tensor) => clone_host!(Bool, tensor),
54 Tensor::C32(tensor) => clone_host!(C32, tensor),
55 Tensor::C64(tensor) => clone_host!(C64, tensor),
56 }
57}
58
59#[cfg(test)]
60fn materialize_tensor_view(
61 buffers: &mut BufferPool,
62 op: &'static str,
63 view: TensorView<'_>,
64) -> Result<Tensor> {
65 macro_rules! materialize {
66 ($variant:ident, $view:expr) => {{
67 Ok(Tensor::$variant(typed_materialize_view_for_tests(
68 buffers, &$view, op,
69 )?))
70 }};
71 }
72 match view {
73 TensorView::F32(view) => materialize!(F32, view),
74 TensorView::F64(view) => materialize!(F64, view),
75 TensorView::I32(view) => materialize!(I32, view),
76 TensorView::I64(view) => materialize!(I64, view),
77 TensorView::Bool(view) => materialize!(Bool, view),
78 TensorView::C32(view) => materialize!(C32, view),
79 TensorView::C64(view) => materialize!(C64, view),
80 }
81}
82
83#[cfg(test)]
84fn typed_materialize_view_for_tests<T, R>(
85 buffers: &mut BufferPool,
86 view: &TypedTensorView<'_, T, R>,
87 op: &'static str,
88) -> Result<TypedTensor<T, R>>
89where
90 T: Copy + Clone + PoolScalar + 'static,
91 R: TensorRank,
92{
93 if view.backend_buffer().is_some() {
94 return Err(cpu_backend_buffer_error(op));
95 }
96 let src: StridedView<'_, T, Identity> = StridedView::new(
97 view.host_storage()?,
98 view.shape(),
99 view.strides(),
100 view.offset(),
101 )
102 .map_err(|err| Error::backend_source(op, err))?;
103 let mut out = PooledUninitOutput::<T>::new(buffers, view.shape().to_vec())?;
104 map_into(&mut out.as_uninit_view_mut()?, &src, |x| {
105 MaybeUninit::new(x)
106 })
107 .map_err(|err| Error::backend_source(op, err))?;
108 let out = unsafe { out.assume_init_as::<R>()? };
110 let shape = R::shape_from_vec(view.shape().to_vec().into())
111 .map_err(|err| Error::backend_source(op, err))?;
112 let mut tensor = TypedTensor::from_vec_col_major(shape, out.into_vec_col_major()?.1)?;
113 tensor.set_placement(view.placement().clone());
114 Ok(tensor)
115}