Skip to main content

tenferro_internal_cpu_kernels/
lib.rs

1#![doc(hidden)]
2
3//! Internal ordinary CPU kernel implementations.
4//!
5//! Shared resource ownership is implemented by `tenferro-cpu-basic`; this crate
6//! owns the ordinary dtype-dispatch kernel family.
7
8pub 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    // SAFETY: the successful map replay writes every logical destination element.
109    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}