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}