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.