pub fn squared_factor_norm<S: ScalarSupport + ?Sized>(
input: &TracedTensor,
support: &S,
) -> Result<TracedTensor, Error>Expand description
The loss the connected program minimizes: the squared norm of the triangular factor.
The loss is ordinary f64 work, so only ScalarSupport::to_f64 stands between the
factorization’s scalar and the objective.
§Examples
struct PassThrough;
impl ScalarSupport for PassThrough {
fn qr(&self, input: &TracedTensor) -> Result<(TracedTensor, TracedTensor), Error> {
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, 2])?;
// The loss is the squared norm of the factor the binding returned, as ordinary f64.
let loss = squared_factor_norm(&input, &PassThrough)?;
assert_eq!(loss.rank, 0);§Errors
Returns the binding’s error when the factorisation or the presentation fails: an unsupported dtype, an invalid shape, or a matrix that is not full column rank.