Skip to main content

tenferro_internal_cpu_kernels/
scalar_ops.rs

1//! Scalar-agnostic entry points over the ordinary CPU numerical bodies.
2//!
3//! The preset scalar types reach the ordinary CPU kernels through the typed
4//! pool, which hands out the destination buffer. An external scalar type cannot
5//! use that pool, so these entry points take a caller-provided destination and
6//! the caller's own arithmetic instead, and run the same elementwise and
7//! reduction bodies the preset path uses.
8//!
9//! Nothing here inspects a dtype tag: the element type and the arithmetic are
10//! the caller's. Support for the preset scalars is therefore not a precondition
11//! for using the ordinary CPU numerical path, and no set-specific numerical
12//! specialization is introduced.
13
14pub mod op;
15
16pub use op::{AddOp, BinaryScalarOp, MulOp, SubOp};
17
18use strided_kernel::{reduce, zip_map2_into, StridedView, StridedViewMut};
19use tenferro_tensor::col_major_strides;
20use tenferro_tensor::{DynRank, Host, TypedTensor};
21
22fn strides_for(shape: &[usize]) -> crate::Result<Vec<isize>> {
23    col_major_strides(shape)
24}
25
26fn require_same_shape(op: &'static str, lhs: &[usize], rhs: &[usize]) -> crate::Result<()> {
27    if lhs == rhs {
28        Ok(())
29    } else {
30        Err(crate::Error::shape_mismatch(op, lhs.to_vec(), rhs.to_vec()))
31    }
32}
33
34/// Apply a named binary operation elementwise into a caller-owned destination.
35///
36/// The destination, the two operands, and the operation are all the caller's.
37/// The traversal is the same `zip_map2_into` body the preset scalar types use
38/// through the typed pool; only the destination's origin differs. The operation
39/// is a type rather than a closure so the instantiation depends on the element
40/// type and the operation, not on the call site.
41///
42/// # Examples
43///
44/// ```rust
45/// use tenferro_internal_cpu_kernels::scalar_ops::{scalar_binary_into, AddOp};
46/// use tenferro_tensor::{DynRank, Host, TypedTensor};
47///
48/// let lhs = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![2], vec![1.0_f64, 2.0])?;
49/// let rhs = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![2], vec![10.0_f64, 20.0])?;
50/// let mut out = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![2], vec![0.0_f64, 0.0])?;
51/// scalar_binary_into::<f64, AddOp>("add", &mut out, &lhs, &rhs)?;
52/// assert_eq!(out.as_slice(), &[11.0, 22.0]);
53/// # Ok::<(), Box<dyn std::error::Error>>(())
54/// ```
55///
56/// # Errors
57///
58/// Returns [`Error::Validation`] when the destination shape does not match the operands',
59/// or [`Error::BackendSource`] when the underlying strided traversal rejects the views.
60pub fn scalar_binary_into<T, Op>(
61    op: &'static str,
62    destination: &mut TypedTensor<T, DynRank, Host>,
63    lhs: &TypedTensor<T, DynRank, Host>,
64    rhs: &TypedTensor<T, DynRank, Host>,
65) -> crate::Result<()>
66where
67    T: Copy + Send + Sync,
68    Op: BinaryScalarOp<T>,
69{
70    require_same_shape(op, destination.shape(), lhs.shape())?;
71    require_same_shape(op, lhs.shape(), rhs.shape())?;
72
73    let strides = strides_for(lhs.shape())?;
74    let destination_shape = destination.shape().to_vec();
75    let mut destination_view: StridedViewMut<'_, T> =
76        StridedViewMut::new(destination.host_data_mut(), &destination_shape, &strides, 0)
77            .map_err(|err| crate::Error::backend_source(op, err))?;
78    let lhs_view: StridedView<'_, T> = StridedView::new(lhs.as_slice(), lhs.shape(), &strides, 0)
79        .map_err(|err| crate::Error::backend_source(op, err))?;
80    let rhs_view: StridedView<'_, T> = StridedView::new(rhs.as_slice(), rhs.shape(), &strides, 0)
81        .map_err(|err| crate::Error::backend_source(op, err))?;
82
83    zip_map2_into(&mut destination_view, &lhs_view, &rhs_view, |a, b| {
84        Op::apply(a, b)
85    })
86    .map_err(|err| crate::Error::backend_source(op, err))
87}
88
89/// Fold every element of a caller-owned tensor with a named associative
90/// operation, starting from `init`.
91///
92/// The accumulation order is the backend's and is not part of the contract. The
93/// arithmetic is the caller's, so an extended-precision scalar keeps its low
94/// components exactly as a preset scalar keeps its own.
95///
96/// # Examples
97///
98/// ```rust
99/// use tenferro_internal_cpu_kernels::scalar_ops::{scalar_fold, AddOp};
100/// use tenferro_tensor::{DynRank, Host, TypedTensor};
101///
102/// let values = TypedTensor::<_, DynRank, Host>::from_host_vec_col_major(vec![3], vec![1.0_f64, 2.0, 3.0])?;
103/// let total = scalar_fold::<f64, AddOp>("sum", &values, 0.0_f64)?;
104/// assert_eq!(total, 6.0);
105/// # Ok::<(), Box<dyn std::error::Error>>(())
106/// ```
107///
108/// # Errors
109///
110/// Returns [`Error::BackendSource`] when the strided reduction rejects the view, which
111/// happens when the source's shape has an invalid stride layout.
112pub fn scalar_fold<T, Op>(
113    op: &'static str,
114    source: &TypedTensor<T, DynRank, Host>,
115    init: T,
116) -> crate::Result<T>
117where
118    T: Copy + Send + Sync,
119    Op: BinaryScalarOp<T>,
120{
121    let strides = strides_for(source.shape())?;
122    let view: StridedView<'_, T> = StridedView::new(source.as_slice(), source.shape(), &strides, 0)
123        .map_err(|err| crate::Error::backend_source(op, err))?;
124    reduce(&view, |element| element, |a, b| Op::apply(a, b), init)
125        .map_err(|err| crate::Error::backend_source(op, err))
126}
127
128#[cfg(test)]
129mod tests;