Skip to main content

tenferro_scalar_consumer_algorithm/
lib.rs

1//! The algorithm role of the external-scalar consumer.
2//!
3//! The algorithm states the capabilities it needs and never names a scalar, a provider,
4//! or a dtype-specific entry point. The application supplies the bindings, so the same
5//! source runs with canonical standard support and with an application-added scalar
6//! contribution.
7
8use tenferro_ad::AdContext;
9use tenferro_runtime::{Error, ErrorPhase, TracedTensor};
10
11/// The scalar and operation capabilities the algorithm requires from its binding.
12///
13/// A binding supplies a factorization of the supported scalar and a way to present a
14/// value of that scalar as an ordinary `f64` tensor, which is what the loss and its
15/// derivative are written against.
16///
17/// # Examples
18///
19/// ```rust
20/// # use tenferro_runtime::{Error, ErrorPhase, TracedTensor};
21/// # use tenferro_scalar_consumer_algorithm::{squared_factor_norm, ScalarSupport};
22/// struct Refusing;
23///
24/// impl ScalarSupport for Refusing {
25///     fn qr(&self, _input: &TracedTensor) -> Result<(TracedTensor, TracedTensor), Error> {
26///         Err(Error::runtime_state("Refusing::qr", ErrorPhase::GraphBuild, "unsupported"))
27///     }
28///     fn to_f64(&self, _input: &TracedTensor) -> Result<TracedTensor, Error> {
29///         Err(Error::runtime_state("Refusing::to_f64", ErrorPhase::GraphBuild, "unsupported"))
30///     }
31/// }
32///
33/// let input = TracedTensor::input_concrete_shape(tenferro_tensor::DType::F64, &[2, 1])?;
34/// // The factorization is the first capability the algorithm needs.
35/// assert!(squared_factor_norm(&input, &Refusing).is_err());
36/// # Ok::<(), Box<dyn std::error::Error>>(())
37/// ```
38pub trait ScalarSupport {
39    /// Factor a matrix into `(Q, R)`, where `R` has a positive diagonal.
40    ///
41    /// # Errors
42    ///
43    /// Returns the binding's error when the input's rank or shape is invalid, when the
44    /// input's dtype is unsupported, or when the matrix is not full column rank.
45    ///
46    /// # Examples
47    ///
48    /// ```rust
49    /// # use tenferro_runtime::{Error, ErrorPhase, TracedTensor};
50    /// # use tenferro_scalar_consumer_algorithm::ScalarSupport;
51    /// struct Identity;
52    ///
53    /// impl ScalarSupport for Identity {
54    ///     fn qr(&self, input: &TracedTensor) -> Result<(TracedTensor, TracedTensor), Error> {
55    ///         // A binding decides what a factorization means for its scalar.
56    ///         Ok((input.clone(), input.clone()))
57    ///     }
58    ///     fn to_f64(&self, input: &TracedTensor) -> Result<TracedTensor, Error> {
59    ///         Ok(input.clone())
60    ///     }
61    /// }
62    ///
63    /// let input = TracedTensor::input_concrete_shape(tenferro_tensor::DType::F64, &[2, 1])?;
64    /// let (_q, r) = Identity.qr(&input)?;
65    /// assert_eq!(r.rank, 2);
66    /// # Ok::<(), Box<dyn std::error::Error>>(())
67    /// ```
68    fn qr(&self, input: &TracedTensor) -> Result<(TracedTensor, TracedTensor), Error>;
69
70    /// Present a supported value as an ordinary `f64` tensor.
71    ///
72    /// # Errors
73    ///
74    /// Returns the binding's error when the value's dtype is unsupported or its shape is
75    /// invalid for the presentation the binding defines.
76    ///
77    /// # Examples
78    ///
79    /// ```rust
80    /// # use tenferro_runtime::{Error, TracedTensor};
81    /// # use tenferro_scalar_consumer_algorithm::ScalarSupport;
82    /// struct AlreadyF64;
83    ///
84    /// impl ScalarSupport for AlreadyF64 {
85    ///     fn qr(&self, input: &TracedTensor) -> Result<(TracedTensor, TracedTensor), Error> {
86    ///         Ok((input.clone(), input.clone()))
87    ///     }
88    ///     fn to_f64(&self, input: &TracedTensor) -> Result<TracedTensor, Error> {
89    ///         // The supported scalar is already ordinary `f64`, so no conversion is needed.
90    ///         Ok(input.clone())
91    ///     }
92    /// }
93    ///
94    /// let input = TracedTensor::input_concrete_shape(tenferro_tensor::DType::F64, &[1])?;
95    /// assert_eq!(AlreadyF64.to_f64(&input)?.dtype(), tenferro_tensor::DType::F64);
96    /// # Ok::<(), Box<dyn std::error::Error>>(())
97    /// ```
98    fn to_f64(&self, input: &TracedTensor) -> Result<TracedTensor, Error>;
99}
100
101/// The loss the connected program minimizes: the squared norm of the triangular factor.
102///
103/// The loss is ordinary `f64` work, so only [`ScalarSupport::to_f64`] stands between the
104/// factorization's scalar and the objective.
105///
106/// # Examples
107///
108/// ```rust
109/// # use tenferro_runtime::{Error, TracedTensor};
110/// # use tenferro_scalar_consumer_algorithm::{squared_factor_norm, ScalarSupport};
111/// struct PassThrough;
112///
113/// impl ScalarSupport for PassThrough {
114///     fn qr(&self, input: &TracedTensor) -> Result<(TracedTensor, TracedTensor), Error> {
115///         Ok((input.clone(), input.clone()))
116///     }
117///     fn to_f64(&self, input: &TracedTensor) -> Result<TracedTensor, Error> {
118///         Ok(input.clone())
119///     }
120/// }
121///
122/// let input = TracedTensor::input_concrete_shape(tenferro_tensor::DType::F64, &[2, 2])?;
123/// // The loss is the squared norm of the factor the binding returned, as ordinary f64.
124/// let loss = squared_factor_norm(&input, &PassThrough)?;
125/// assert_eq!(loss.rank, 0);
126/// # Ok::<(), Box<dyn std::error::Error>>(())
127/// ```
128///
129/// # Errors
130///
131/// Returns the binding's error when the factorisation or the presentation fails: an
132/// unsupported dtype, an invalid shape, or a matrix that is not full column rank.
133pub fn squared_factor_norm<S: ScalarSupport + ?Sized>(
134    input: &TracedTensor,
135    support: &S,
136) -> Result<TracedTensor, Error> {
137    let (_q, r) = support.qr(input)?;
138    let narrowed = support.to_f64(&r)?;
139    let squared = narrowed.mul(&narrowed)?;
140    squared.reduce_sum(None)
141}
142
143/// The objective's gradient with respect to the program input.
144///
145/// # Examples
146///
147/// ```rust
148/// # use tenferro_runtime::{Error, ErrorPhase, TracedTensor};
149/// # use tenferro_scalar_consumer_algorithm::{factor_norm_gradient, ScalarSupport};
150/// struct Refusing;
151///
152/// impl ScalarSupport for Refusing {
153///     fn qr(&self, _input: &TracedTensor) -> Result<(TracedTensor, TracedTensor), Error> {
154///         Err(Error::runtime_state("Refusing::qr", ErrorPhase::GraphBuild, "unsupported"))
155///     }
156///     fn to_f64(&self, _input: &TracedTensor) -> Result<TracedTensor, Error> {
157///         Err(Error::runtime_state("Refusing::to_f64", ErrorPhase::GraphBuild, "unsupported"))
158///     }
159/// }
160///
161/// let ad = tenferro_ad::AdContext::builder().build()?;
162/// let input = TracedTensor::input_concrete_shape(tenferro_tensor::DType::F64, &[2, 1])?;
163/// let seed = TracedTensor::input_concrete_shape(tenferro_tensor::DType::F64, &[])?;
164/// // A binding that cannot factor stops the program before the reverse pass.
165/// assert!(factor_norm_gradient(&ad, &input, &seed, &Refusing).is_err());
166/// # Ok::<(), Box<dyn std::error::Error>>(())
167/// ```
168///
169/// # Errors
170///
171/// Returns the binding's error when the factorisation or the presentation fails with an
172/// unsupported dtype or an invalid shape, [`Error::RuntimeStateSource`] when the input does
173/// not reach the loss, or the reverse pass error when the program cannot be built.
174pub fn factor_norm_gradient<S: ScalarSupport + ?Sized>(
175    ad: &AdContext,
176    input: &TracedTensor,
177    seed: &TracedTensor,
178    support: &S,
179) -> Result<TracedTensor, Error> {
180    let loss = squared_factor_norm(input, support)?;
181    let gradients = ad.vjp_many(&loss, &[input], seed)?;
182    gradients.into_iter().next().flatten().ok_or_else(|| {
183        Error::runtime_state(
184            "factor_norm_gradient",
185            ErrorPhase::GraphBuild,
186            "the input is inactive",
187        )
188    })
189}
190
191#[cfg(test)]
192mod tests;