Skip to main content

AdContext

Struct AdContext 

Source
pub struct AdContext { /* private fields */ }
Expand description

Explicit automatic-differentiation context.

AdContext owns the extension AD rules used by traced AD transforms. It also owns the AD transform cache shared by context-driven traced AD and eager runtimes created from this context.

§Examples

use tenferro_ad::AdContext;

let ad = AdContext::builder().build().unwrap();
assert!(ad
    .semantic_extension_rules()
    .lookup_linearize("example.missing.v1")
    .is_none());

Implementations§

Source§

impl AdContext

Source

pub fn builder() -> AdContextBuilder

Start building an explicit AD context.

§Examples
use tenferro_ad::AdContext;

let _builder = AdContext::builder();
Source

pub fn semantic_extension_rules(&self) -> &SemanticExtensionRuleSet

Return semantic-program extension AD rules owned by this context.

§Examples
use tenferro_ad::AdContext;

let ad = AdContext::builder().build().unwrap();
assert!(ad
    .semantic_extension_rules()
    .lookup_linearize("example.missing.v1")
    .is_none());
Source

pub fn jvp_program( &self, input: &FrozenProgram, active_inputs: &[bool], ) -> Result<SemanticAdProgram, SemanticAdTransformError>

Transform a frozen semantic program into its forward-mode derivative.

active_inputs follows source-program input order. Active tangent seeds are appended after all primal inputs.

§Errors

Returns SemanticAdTransformError::ActivityArity when active_inputs has the wrong length, SemanticAdTransformError::Extension when an extension rule rejects the transform, or the corresponding Query, Build, Finish, or Cache variant when program import, construction, finalization, or cache access fails.

Source

pub fn vjp_program( &self, input: &FrozenProgram, active_inputs: &[bool], active_outputs: &[bool], ) -> Result<SemanticAdProgram, SemanticAdTransformError>

Transform a frozen semantic program into its reverse-mode derivative.

active_inputs selects requested primal-input cotangents and active_outputs selects primal outputs that receive appended seeds.

§Errors

Returns SemanticAdTransformError::ActivityArity when either activity mask has the wrong length, SemanticAdTransformError::Extension when an extension rule rejects the transform, or the corresponding Query, Build, Finish, or Cache variant when program import, construction, finalization, or cache access fails.

Source

pub fn ad_transform_cache_limits(&self) -> Result<AdTransformCacheLimits>

Return AD transform cache retention limits.

§Examples
use tenferro_ad::AdContext;

let ad = AdContext::builder().build().unwrap();
assert!(ad.ad_transform_cache_limits().unwrap().max_entries().get() > 0);
§Errors

Returns tenferro_runtime::Error::RuntimeState if the cache lock is poisoned or its state cannot be inspected.

Source

pub fn set_ad_transform_cache_limits( &self, limits: AdTransformCacheLimits, ) -> Result<()>

Replace AD transform cache retention limits.

§Examples
use std::num::NonZeroUsize;
use tenferro_ad::{AdContext, AdTransformCacheLimits};

let ad = AdContext::builder().build().unwrap();
let limits = AdTransformCacheLimits::new(NonZeroUsize::new(1).unwrap());
ad.set_ad_transform_cache_limits(limits).unwrap();
assert_eq!(ad.ad_transform_cache_limits().unwrap(), limits);
§Errors

Returns tenferro_runtime::Error::RuntimeState if the cache lock is poisoned while updating the limits.

Source

pub fn clear_ad_transform_caches(&self) -> Result<()>

Clear AD transform cache entries owned by this context.

§Examples
use tenferro_ad::AdContext;

let ad = AdContext::builder().build().unwrap();
ad.clear_ad_transform_caches().unwrap();
assert_eq!(ad.ad_transform_cache_stats().unwrap().entries, 0);
§Errors

Returns tenferro_runtime::Error::RuntimeState if the cache lock is poisoned while clearing entries.

Source

pub fn ad_transform_cache_stats(&self) -> Result<CacheStats>

Return AD transform cache-entry and retained-byte stats.

§Examples
use tenferro_ad::AdContext;

let ad = AdContext::builder().build().unwrap();
assert_eq!(ad.ad_transform_cache_stats().unwrap().entries, 0);
§Errors

Returns tenferro_runtime::Error::RuntimeState if the cache lock is poisoned while collecting statistics.

Source

pub fn clear_caches(&self) -> Result<()>

Clear every cache owned by this AD context.

§Examples
use tenferro_ad::AdContext;

let ad = AdContext::builder().build().unwrap();
ad.clear_caches().unwrap();
assert_eq!(ad.cache_stats().unwrap().ad_transforms.entries, 0);
§Errors

Returns tenferro_runtime::Error::RuntimeState if either owned cache cannot be locked because its state is poisoned.

Source

pub fn cache_stats(&self) -> Result<AdContextCacheStats>

Return aggregate cache-entry and retained-byte stats for this AD context.

§Examples
use tenferro_ad::AdContext;

let ad = AdContext::builder().build().unwrap();
assert_eq!(ad.cache_stats().unwrap().ad_transforms.retained_bytes, 0);
§Errors

Returns tenferro_runtime::Error::RuntimeState if an owned cache lock is poisoned while collecting statistics.

Source

pub fn grad( &self, output: &TracedTensor, wrt: &TracedTensor, ) -> Result<TracedTensor>

Gradient of a scalar traced output with respect to a traced input.

For complex scalar outputs, tenferro returns the Hermitian-adjoint cotangent. To compare seed-1 scalar gradients with JAX’s public grad values, use the complex conjugate of this result. See https://tensor4all.org/tenferro-rs/guides/complex-ad.html.

§Examples
use tenferro_ad::AdContext;
use tenferro_runtime::TracedTensor;

let ad = AdContext::builder().build().unwrap();
let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
let loss = (&x * &x).unwrap();
let grad = ad.grad(&loss, &x).unwrap();
assert_eq!(grad.rank, 0);
§Errors

Returns tenferro_runtime::Error::NonScalarGrad when output is not scalar, tenferro_runtime::Error::UnsupportedAdRule when a graph op lacks a registered rule, or a typed tenferro_runtime::Error::Validation / backend error when graph metadata or execution is invalid.

Source

pub fn grad_optional( &self, output: &TracedTensor, wrt: &TracedTensor, ) -> Result<Option<TracedTensor>>

Gradient that returns None when wrt is inactive.

§Examples
use tenferro_ad::AdContext;
use tenferro_runtime::TracedTensor;

let ad = AdContext::builder().build().unwrap();
let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
let loss = (&x * &x).unwrap();
assert!(ad.grad_optional(&loss, &x).unwrap().is_some());
§Errors

Returns tenferro_runtime::Error::NonScalarGrad for a non-scalar output, tenferro_runtime::Error::UnsupportedAdRule for an unregistered AD rule, or a typed tenferro_runtime::Error::Validation / backend error from graph construction and execution.

Source

pub fn jvp( &self, output: &TracedTensor, wrt: &TracedTensor, tangent: &TracedTensor, ) -> Result<TracedTensor>

Forward-mode Jacobian-vector product.

§Examples
use tenferro_ad::AdContext;
use tenferro_runtime::TracedTensor;

let ad = AdContext::builder().build().unwrap();
let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
let dx = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
let y = (&x * &x).unwrap();
let dy = ad.jvp(&y, &x, &dx).unwrap();
assert_eq!(dy.rank, 0);
§Errors

Returns tenferro_runtime::Error::UnsupportedAdRule when the graph has no JVP rule, tenferro_runtime::Error::Validation for inconsistent tangent metadata, or a typed backend/runtime-state error during evaluation.

Source

pub fn jvp_optional( &self, output: &TracedTensor, wrt: &TracedTensor, tangent: &TracedTensor, ) -> Result<Option<TracedTensor>>

Forward-mode Jacobian-vector product that returns None for inactive output.

§Examples
use tenferro_ad::AdContext;
use tenferro_runtime::TracedTensor;

let ad = AdContext::builder().build().unwrap();
let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
let dx = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
let y = (&x * &x).unwrap();
assert!(ad.jvp_optional(&y, &x, &dx).unwrap().is_some());
§Errors

Returns tenferro_runtime::Error::UnsupportedAdRule when the graph has no JVP rule, tenferro_runtime::Error::Validation for inconsistent tangent metadata, or a typed backend/runtime-state error during evaluation.

Source

pub fn jvp_many( &self, output: &TracedTensor, wrt_tangents: &[(&TracedTensor, &TracedTensor)], ) -> Result<Option<TracedTensor>>

Forward-mode directional derivative for multiple distinct traced leaves.

Reachable leaves are transformed together in one derivative graph. Unreachable leaves contribute nothing; an empty or fully unreachable request returns None. Duplicate wrt leaves are rejected before the transform because one semantic seed slot cannot accept two tangents.

§Examples
use tenferro_ad::AdContext;
use tenferro_runtime::TracedTensor;

let ad = AdContext::builder().build().unwrap();
let x = TracedTensor::from_vec_col_major(vec![], vec![2.0_f64]).unwrap();
let y = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
let dx = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
let dy = TracedTensor::from_vec_col_major(vec![], vec![4.0_f64]).unwrap();
let output = (&x * &y).unwrap();
assert!(ad.jvp_many(&output, &[(&x, &dx), (&y, &dy)]).unwrap().is_some());
§Errors

Returns tenferro_runtime::Error::Validation for duplicate leaves or incompatible tangent metadata, tenferro_runtime::Error::UnsupportedAdRule when a required rule is unavailable, or a typed runtime-state error when derivative graph construction fails.

§Deferred errors

Symbolic shape constraints may fail during later compilation or execution.

Source

pub fn vjp( &self, output: &TracedTensor, wrt: &TracedTensor, cotangent: &TracedTensor, ) -> Result<TracedTensor>

Reverse-mode vector-Jacobian product.

Complex cotangents use tenferro’s Hermitian real-inner-product convention. Non-real complex cotangent seeds therefore need an explicit seed-convention comparison when matching JAX. See https://tensor4all.org/tenferro-rs/guides/complex-ad.html.

§Examples
use tenferro_ad::AdContext;
use tenferro_runtime::TracedTensor;

let ad = AdContext::builder().build().unwrap();
let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
let dy = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
let y = (&x * &x).unwrap();
let dx = ad.vjp(&y, &x, &dy).unwrap();
assert_eq!(dx.rank, 0);
§Errors

Returns tenferro_runtime::Error::Validation when the cotangent metadata is incompatible, tenferro_runtime::Error::UnsupportedAdRule when a VJP rule is unavailable, or a typed backend/runtime-state error during execution.

Source

pub fn vjp_optional( &self, output: &TracedTensor, wrt: &TracedTensor, cotangent: &TracedTensor, ) -> Result<Option<TracedTensor>>

Reverse-mode vector-Jacobian product that returns None for inactive input.

§Examples
use tenferro_ad::AdContext;
use tenferro_runtime::TracedTensor;

let ad = AdContext::builder().build().unwrap();
let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
let dy = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
let y = (&x * &x).unwrap();
assert!(ad.vjp_optional(&y, &x, &dy).unwrap().is_some());
§Errors

Returns tenferro_runtime::Error::Validation when the cotangent metadata is incompatible, tenferro_runtime::Error::UnsupportedAdRule when a VJP rule is unavailable, or a typed backend/runtime-state error during execution.

Source

pub fn vjp_many( &self, output: &TracedTensor, wrts: &[&TracedTensor], cotangent: &TracedTensor, ) -> Result<Vec<Option<TracedTensor>>>

Reverse-mode products for multiple traced leaves in one derivative graph.

Results align with wrts; unreachable leaves produce None. Duplicate leaves are allowed and repeat the same traced derivative without accumulating the cotangent twice. An empty request validates that the cotangent has concrete data, then returns an empty vector.

§Examples
use tenferro_ad::AdContext;
use tenferro_runtime::TracedTensor;

let ad = AdContext::builder().build().unwrap();
let x = TracedTensor::from_vec_col_major(vec![], vec![2.0_f64]).unwrap();
let y = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
let seed = TracedTensor::from_vec_col_major(vec![], vec![1.0_f64]).unwrap();
let output = (&x * &y).unwrap();
let products = ad.vjp_many(&output, &[&x, &y], &seed).unwrap();
assert!(products.iter().all(Option::is_some));
§Errors

Returns tenferro_runtime::Error::Validation for invalid cotangent metadata, tenferro_runtime::Error::UnsupportedAdRule when a required rule is unavailable, or a typed runtime-state error when derivative graph construction fails.

§Deferred errors

Symbolic shape constraints may fail during later compilation or execution.

Trait Implementations§

Source§

impl Clone for AdContext

Source§

fn clone(&self) -> AdContext

Returns a duplicate of the value. Read more
1.0.0 (const: unstable) · Source§

fn clone_from(&mut self, source: &Self)

Performs copy-assignment from source. Read more
Source§

impl Debug for AdContext

Source§

fn fmt(&self, f: &mut Formatter<'_>) -> Result

Formats the value using the given formatter. Read more

Auto Trait Implementations§

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
§

impl<T> ByRef<T> for T

§

fn by_ref(&self) -> &T

Source§

impl<T> CloneToUninit for T
where T: Clone,

Source§

unsafe fn clone_to_uninit(&self, dest: *mut u8)

🔬This is a nightly-only experimental API. (clone_to_uninit)
Performs copy-assignment from self to dest. Read more
Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

§

impl<T, U> Imply<T> for U
where T: ?Sized, U: ?Sized,

Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> IntoEither for T

Source§

fn into_either(self, into_left: bool) -> Either<Self, Self>

Converts self into a Left variant of Either<Self, Self> if into_left is true. Converts self into a Right variant of Either<Self, Self> otherwise. Read more
Source§

fn into_either_with<F>(self, into_left: F) -> Either<Self, Self>
where F: FnOnce(&Self) -> bool,

Converts self into a Left variant of Either<Self, Self> if into_left(&self) returns true. Converts self into a Right variant of Either<Self, Self> otherwise. Read more
§

impl<T> MaybeSend for T
where T: Send,

§

impl<T> MaybeSendSync for T
where T: Send + Sync,

§

impl<T> MaybeSync for T
where T: Sync,

§

impl<T> Pointable for T

§

const ALIGN: usize

The alignment of pointer.
§

type Init = T

The type for initializers.
§

unsafe fn init(init: <T as Pointable>::Init) -> usize

Initializes a with the given initializer. Read more
§

unsafe fn deref<'a>(ptr: usize) -> &'a T

Dereferences the given pointer. Read more
§

unsafe fn deref_mut<'a>(ptr: usize) -> &'a mut T

Mutably dereferences the given pointer. Read more
§

unsafe fn drop(ptr: usize)

Drops the object pointed to by the given pointer. Read more
Source§

impl<T> ToOwned for T
where T: Clone,

Source§

type Owned = T

The resulting type after obtaining ownership.
Source§

fn to_owned(&self) -> T

Creates owned data from borrowed data, usually by cloning. Read more
Source§

fn clone_into(&self, target: &mut T)

Uses borrowed data to replace owned data, usually by cloning. Read more
Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = Infallible

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, <T as TryFrom<U>>::Error>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.