pub trait ScalarSupport {
// Required methods
fn qr(
&self,
input: &TracedTensor,
) -> Result<(TracedTensor, TracedTensor), Error>;
fn to_f64(&self, input: &TracedTensor) -> Result<TracedTensor, Error>;
}Expand description
The scalar and operation capabilities the algorithm requires from its binding.
A binding supplies a factorization of the supported scalar and a way to present a
value of that scalar as an ordinary f64 tensor, which is what the loss and its
derivative are written against.
§Examples
struct Refusing;
impl ScalarSupport for Refusing {
fn qr(&self, _input: &TracedTensor) -> Result<(TracedTensor, TracedTensor), Error> {
Err(Error::runtime_state("Refusing::qr", ErrorPhase::GraphBuild, "unsupported"))
}
fn to_f64(&self, _input: &TracedTensor) -> Result<TracedTensor, Error> {
Err(Error::runtime_state("Refusing::to_f64", ErrorPhase::GraphBuild, "unsupported"))
}
}
let input = TracedTensor::input_concrete_shape(tenferro_tensor::DType::F64, &[2, 1])?;
// The factorization is the first capability the algorithm needs.
assert!(squared_factor_norm(&input, &Refusing).is_err());Required Methods§
Sourcefn qr(
&self,
input: &TracedTensor,
) -> Result<(TracedTensor, TracedTensor), Error>
fn qr( &self, input: &TracedTensor, ) -> Result<(TracedTensor, TracedTensor), Error>
Factor a matrix into (Q, R), where R has a positive diagonal.
§Errors
Returns the binding’s error when the input’s rank or shape is invalid, when the input’s dtype is unsupported, or when the matrix is not full column rank.
§Examples
struct Identity;
impl ScalarSupport for Identity {
fn qr(&self, input: &TracedTensor) -> Result<(TracedTensor, TracedTensor), Error> {
// A binding decides what a factorization means for its scalar.
Ok((input.clone(), input.clone()))
}
fn to_f64(&self, input: &TracedTensor) -> Result<TracedTensor, Error> {
Ok(input.clone())
}
}
let input = TracedTensor::input_concrete_shape(tenferro_tensor::DType::F64, &[2, 1])?;
let (_q, r) = Identity.qr(&input)?;
assert_eq!(r.rank, 2);Sourcefn to_f64(&self, input: &TracedTensor) -> Result<TracedTensor, Error>
fn to_f64(&self, input: &TracedTensor) -> Result<TracedTensor, Error>
Present a supported value as an ordinary f64 tensor.
§Errors
Returns the binding’s error when the value’s dtype is unsupported or its shape is invalid for the presentation the binding defines.
§Examples
struct AlreadyF64;
impl ScalarSupport for AlreadyF64 {
fn qr(&self, input: &TracedTensor) -> Result<(TracedTensor, TracedTensor), Error> {
Ok((input.clone(), input.clone()))
}
fn to_f64(&self, input: &TracedTensor) -> Result<TracedTensor, Error> {
// The supported scalar is already ordinary `f64`, so no conversion is needed.
Ok(input.clone())
}
}
let input = TracedTensor::input_concrete_shape(tenferro_tensor::DType::F64, &[1])?;
assert_eq!(AlreadyF64.to_f64(&input)?.dtype(), tenferro_tensor::DType::F64);Dyn Compatibility§
This trait is dyn compatible.
In older versions of Rust, dyn compatibility was called "object safety".