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
8/// The Rust scalar type behind a preset variant name a macro received.
9#[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
46// Re-exported so the exported dispatch macros can name them with `$crate` paths.
47pub 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        // A caller-owned payload must be duplicated by its owner, because this
109        // crate cannot clone an erased element type.
110        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)]
119/// The typed tensor behind `tensor`, or this module's refusal for a dtype it cannot clone.
120///
121/// Callers reach this from a match on `tensor.dtype()`, so `None` means the tag table and the
122/// runtime dtype disagree rather than a caller mistake; the refusal carries the same text the
123/// externally defined arm uses.
124fn 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    // SAFETY: the successful map replay writes every logical destination element.
187    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;