Skip to main content

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;