Skip to main content

tenferro_internal_cpu_kernels/scalar_ops/
op.rs

1//! Scalar operations named by type.
2//!
3//! An operation passed to the scalar-agnostic entry points is a *type*, not a
4//! closure. A closure is part of a generic function's type parameters, so every
5//! call site would produce its own instantiation of the whole `strided-kernel`
6//! body. A named operation keeps the instantiation keyed on the element type and
7//! the operation, so two scalar sets that call the same operation share one
8//! compiled kernel instead of one per set.
9
10/// A binary operation between two scalars of the same type.
11///
12/// Implement this for a marker type in the crate that owns the operation. That
13/// is how an external scalar's contribution supplies its own arithmetic.
14///
15/// # Examples
16///
17/// ```rust
18/// use tenferro_internal_cpu_kernels::scalar_ops::{AddOp, BinaryScalarOp};
19///
20/// assert_eq!(<AddOp as BinaryScalarOp<f64>>::apply(1.0, 2.0), 3.0);
21/// // Integers wrap, so the shared path agrees with the preset path in every build.
22/// assert_eq!(<AddOp as BinaryScalarOp<i32>>::apply(i32::MAX, 1), i32::MIN);
23/// ```
24pub trait BinaryScalarOp<T> {
25    /// Apply the operation.
26    ///
27    /// # Examples
28    ///
29    /// ```rust
30    /// use tenferro_internal_cpu_kernels::scalar_ops::{BinaryScalarOp, MulOp};
31    ///
32    /// assert_eq!(<MulOp as BinaryScalarOp<i64>>::apply(3, 4), 12);
33    /// ```
34    fn apply(lhs: T, rhs: T) -> T;
35}
36
37/// Addition.
38///
39/// # Examples
40///
41/// ```rust
42/// use tenferro_internal_cpu_kernels::scalar_ops::{AddOp, BinaryScalarOp};
43///
44/// assert_eq!(<AddOp as BinaryScalarOp<i32>>::apply(1, 2), 3);
45/// ```
46pub struct AddOp;
47
48/// Subtraction.
49///
50/// # Examples
51///
52/// ```rust
53/// use tenferro_internal_cpu_kernels::scalar_ops::{BinaryScalarOp, SubOp};
54///
55/// assert_eq!(<SubOp as BinaryScalarOp<f64>>::apply(5.0, 2.0), 3.0);
56/// ```
57pub struct SubOp;
58
59/// Multiplication.
60///
61/// # Examples
62///
63/// ```rust
64/// use tenferro_internal_cpu_kernels::scalar_ops::{BinaryScalarOp, MulOp};
65///
66/// assert_eq!(<MulOp as BinaryScalarOp<f64>>::apply(3.0, 4.0), 12.0);
67/// ```
68pub struct MulOp;
69
70impl<T> BinaryScalarOp<T> for AddOp
71where
72    T: tenferro_tensor_core::ScalarArithmetic,
73{
74    /// The scalar contract's addition, which wraps for the integer members.
75    ///
76    /// Going through the contract rather than the operator matters for integers: the operator
77    /// panics on overflow in a debug build and wraps in a release build, while the contract
78    /// wraps in both, which is what the preset path and `ScalarArithmetic::scalar_add` promise.
79    fn apply(lhs: T, rhs: T) -> T {
80        tenferro_tensor_core::ScalarArithmetic::scalar_add(lhs, rhs)
81    }
82}
83
84impl<T> BinaryScalarOp<T> for SubOp
85where
86    T: tenferro_tensor_core::ScalarArithmetic,
87{
88    /// The scalar contract's subtraction, which wraps for the integer members.
89    fn apply(lhs: T, rhs: T) -> T {
90        tenferro_tensor_core::ScalarArithmetic::scalar_sub(lhs, rhs)
91    }
92}
93
94impl<T> BinaryScalarOp<T> for MulOp
95where
96    T: tenferro_tensor_core::ScalarArithmetic,
97{
98    /// The scalar contract's product, which wraps for the integer members.
99    fn apply(lhs: T, rhs: T) -> T {
100        tenferro_tensor_core::ScalarArithmetic::scalar_mul(lhs, rhs)
101    }
102}