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;