Skip to main content

factor_norm_gradient

Function factor_norm_gradient 

Source
pub fn factor_norm_gradient<S: ScalarSupport + ?Sized>(
    ad: &AdContext,
    input: &TracedTensor,
    seed: &TracedTensor,
    support: &S,
) -> Result<TracedTensor, Error>
Expand description

The objective’s gradient with respect to the program input.

§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 ad = tenferro_ad::AdContext::builder().build()?;
let input = TracedTensor::input_concrete_shape(tenferro_tensor::DType::F64, &[2, 1])?;
let seed = TracedTensor::input_concrete_shape(tenferro_tensor::DType::F64, &[])?;
// A binding that cannot factor stops the program before the reverse pass.
assert!(factor_norm_gradient(&ad, &input, &seed, &Refusing).is_err());

§Errors

Returns the binding’s error when the factorisation or the presentation fails with an unsupported dtype or an invalid shape, Error::RuntimeStateSource when the input does not reach the loss, or the reverse pass error when the program cannot be built.