Skip to main content

squared_factor_norm

Function squared_factor_norm 

Source
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.