tenferro_tensor_core/scalar.rs
1//! Open scalar contracts for the host tensor data model.
2//!
3//! The preset scalar types are ordinary members of these contracts, not special
4//! cases: the same table that enumerates them for the tag also declares their
5//! scalar properties. A downstream crate implements [`Scalar`] for its own type
6//! to take part in the same machinery.
7//!
8//! Scalars are classified by the algebra their ordinary arithmetic belongs to.
9//! The host tensor data model stores values of any [`Scalar`]; differentiation
10//! is a separate question answered by [`ad_admission`].
11
12/// Algebra that a scalar's ordinary arithmetic belongs to.
13///
14/// # Examples
15///
16/// ```rust
17/// use tenferro_tensor_core::{Scalar, ScalarDomain};
18///
19/// assert_eq!(<f64 as Scalar>::DOMAIN, ScalarDomain::Field);
20/// assert_eq!(<bool as Scalar>::DOMAIN, ScalarDomain::NonField);
21/// ```
22#[derive(Clone, Copy, Debug, PartialEq, Eq)]
23#[non_exhaustive]
24pub enum ScalarDomain {
25 /// Addition, subtraction, multiplication, and multiplication by a negative
26 /// value follow the ordinary real or complex field rules, so the canonical
27 /// mathematical derivative definitions apply.
28 Field,
29 /// Any other algebra: tropical or min-plus, boolean, saturating, or another
30 /// semiring. Such a scalar can still be stored and computed with, but the
31 /// canonical field derivative rules are not valid for it.
32 NonField,
33}
34
35/// A scalar the host tensor data model can store and move.
36///
37/// This contract carries only storage and representation properties. It does
38/// not require arithmetic, a dtype tag, or any operation support, so declaring
39/// a scalar never forces unrelated implementations.
40///
41/// # Examples
42///
43/// ```rust
44/// use tenferro_tensor_core::{Scalar, ScalarDomain};
45///
46/// fn domain_of<T: Scalar>() -> ScalarDomain {
47/// T::DOMAIN
48/// }
49///
50/// assert_eq!(domain_of::<i32>(), ScalarDomain::Field);
51/// ```
52pub trait Scalar: Copy + Send + Sync + 'static {
53 /// Algebra of this scalar's ordinary arithmetic.
54 const DOMAIN: ScalarDomain;
55}
56
57/// Arithmetic a scalar supports under its own rules.
58///
59/// The operations use the scalar's own semantics, so an integer implementation
60/// wraps exactly as the existing integer kernels do. Complex scalars qualify as
61/// [`ScalarDomain::Field`]. `bool` does not implement this trait: it is a
62/// storable scalar without arithmetic.
63///
64/// # Examples
65///
66/// ```rust
67/// use tenferro_tensor_core::ScalarArithmetic;
68///
69/// assert_eq!(<f64 as ScalarArithmetic>::scalar_add(1.0, 2.0), 3.0);
70/// assert_eq!(<i32 as ScalarArithmetic>::scalar_add(i32::MAX, 1), i32::MIN);
71/// ```
72pub trait ScalarArithmetic: Scalar {
73 /// Additive identity.
74 ///
75 /// # Examples
76 ///
77 /// ```rust
78 /// use tenferro_tensor_core::ScalarArithmetic;
79 ///
80 /// assert_eq!(<f64 as ScalarArithmetic>::scalar_zero(), 0.0);
81 /// ```
82 fn scalar_zero() -> Self;
83
84 /// Multiplicative identity.
85 ///
86 /// # Examples
87 ///
88 /// ```rust
89 /// use tenferro_tensor_core::ScalarArithmetic;
90 ///
91 /// assert_eq!(<f64 as ScalarArithmetic>::scalar_one(), 1.0);
92 /// ```
93 fn scalar_one() -> Self;
94
95 /// Sum under this scalar's own semantics.
96 ///
97 /// # Examples
98 ///
99 /// ```rust
100 /// use tenferro_tensor_core::ScalarArithmetic;
101 ///
102 /// assert_eq!(<f64 as ScalarArithmetic>::scalar_add(1.0, 2.0), 3.0);
103 /// ```
104 fn scalar_add(self, rhs: Self) -> Self;
105
106 /// Difference under this scalar's own semantics.
107 ///
108 /// # Examples
109 ///
110 /// ```rust
111 /// use tenferro_tensor_core::ScalarArithmetic;
112 ///
113 /// assert_eq!(<f64 as ScalarArithmetic>::scalar_sub(1.0, 2.0), -1.0);
114 /// ```
115 fn scalar_sub(self, rhs: Self) -> Self;
116
117 /// Product under this scalar's own semantics.
118 ///
119 /// # Examples
120 ///
121 /// ```rust
122 /// use tenferro_tensor_core::ScalarArithmetic;
123 ///
124 /// assert_eq!(<f64 as ScalarArithmetic>::scalar_mul(3.0, 4.0), 12.0);
125 /// ```
126 fn scalar_mul(self, rhs: Self) -> Self;
127}
128
129/// Why a scalar may not be differentiated at the requested order.
130///
131/// This is the query result of [`ad_admission`]; it never represents a gradient.
132///
133/// # Examples
134///
135/// ```rust
136/// use tenferro_tensor_core::{ad_admission, AdAdmissionError};
137///
138/// assert_eq!(
139/// ad_admission::<f64>(2),
140/// Err(AdAdmissionError::UnsupportedAdOrder { order: 2 })
141/// );
142/// ```
143#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
144#[non_exhaustive]
145pub enum AdAdmissionError {
146 /// Only first-order differentiation is admitted for a non-preset scalar.
147 #[error("differentiation order {order} is not supported")]
148 UnsupportedAdOrder {
149 /// Requested derivative order.
150 order: u32,
151 },
152 /// The scalar's arithmetic is not a field, so the canonical field
153 /// derivative rules do not apply.
154 #[error("the scalar's arithmetic is not an ordinary field")]
155 NonFieldScalar,
156 /// The scalar is admissible in principle, but no derivative rules exist for
157 /// it in this build.
158 #[error("no derivative rules are available for this scalar")]
159 AdRuleUnavailable,
160}
161
162/// Answer whether the shared differentiation paths may differentiate `T` at
163/// `order`.
164///
165/// This is a query. It does not run kernels, register rules, or change existing
166/// differentiation, and a rejection is always explicit: an unsupported scalar
167/// or order is never silently treated as a zero gradient.
168///
169/// # Examples
170///
171/// ```rust
172/// use tenferro_tensor_core::{ad_admission, AdAdmissionError};
173///
174/// // A non-field scalar is rejected even at first order.
175/// assert_eq!(ad_admission::<bool>(1), Err(AdAdmissionError::NonFieldScalar));
176///
177/// // A field scalar with no rules yet is rejected as unavailable, not as zero.
178/// assert_eq!(
179/// ad_admission::<f64>(1),
180/// Err(AdAdmissionError::AdRuleUnavailable)
181/// );
182/// ```
183///
184/// # Errors
185///
186/// Returns [`AdAdmissionError::UnsupportedAdOrder`] for any order other than
187/// one, [`AdAdmissionError::NonFieldScalar`] for a scalar whose arithmetic is
188/// not an ordinary field, and [`AdAdmissionError::AdRuleUnavailable`] when no
189/// rules exist for an otherwise admissible scalar.
190pub fn ad_admission<T: Scalar>(order: u32) -> Result<(), AdAdmissionError> {
191 if order != 1 {
192 return Err(AdAdmissionError::UnsupportedAdOrder { order });
193 }
194 match T::DOMAIN {
195 ScalarDomain::Field => Err(AdAdmissionError::AdRuleUnavailable),
196 ScalarDomain::NonField => Err(AdAdmissionError::NonFieldScalar),
197 }
198}
199
200#[cfg(test)]
201mod tests;