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
impl AdContext
Sourcepub fn builder() -> AdContextBuilder
pub fn builder() -> AdContextBuilder
Start building an explicit AD context.
§Examples
use tenferro_ad::AdContext;
let _builder = AdContext::builder();Sourcepub fn semantic_extension_rules(&self) -> &SemanticExtensionRuleSet
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());Sourcepub fn jvp_program(
&self,
input: &FrozenProgram,
active_inputs: &[bool],
) -> Result<SemanticAdProgram, SemanticAdTransformError>
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.
Sourcepub fn vjp_program(
&self,
input: &FrozenProgram,
active_inputs: &[bool],
active_outputs: &[bool],
) -> Result<SemanticAdProgram, SemanticAdTransformError>
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.
Sourcepub fn ad_transform_cache_limits(&self) -> Result<AdTransformCacheLimits>
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.
Sourcepub fn set_ad_transform_cache_limits(
&self,
limits: AdTransformCacheLimits,
) -> Result<()>
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.
Sourcepub fn clear_ad_transform_caches(&self) -> Result<()>
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.
Sourcepub fn ad_transform_cache_stats(&self) -> Result<CacheStats>
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.
Sourcepub fn clear_caches(&self) -> Result<()>
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.
Sourcepub fn cache_stats(&self) -> Result<AdContextCacheStats>
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.
Sourcepub fn grad(
&self,
output: &TracedTensor,
wrt: &TracedTensor,
) -> Result<TracedTensor>
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.
Sourcepub fn grad_optional(
&self,
output: &TracedTensor,
wrt: &TracedTensor,
) -> Result<Option<TracedTensor>>
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.
Sourcepub fn jvp(
&self,
output: &TracedTensor,
wrt: &TracedTensor,
tangent: &TracedTensor,
) -> Result<TracedTensor>
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.
Sourcepub fn jvp_optional(
&self,
output: &TracedTensor,
wrt: &TracedTensor,
tangent: &TracedTensor,
) -> Result<Option<TracedTensor>>
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.
Sourcepub fn vjp(
&self,
output: &TracedTensor,
wrt: &TracedTensor,
cotangent: &TracedTensor,
) -> Result<TracedTensor>
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.
Sourcepub fn vjp_optional(
&self,
output: &TracedTensor,
wrt: &TracedTensor,
cotangent: &TracedTensor,
) -> Result<Option<TracedTensor>>
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.
Trait Implementations§
Auto Trait Implementations§
impl Freeze for AdContext
impl !RefUnwindSafe for AdContext
impl Send for AdContext
impl Sync for AdContext
impl Unpin for AdContext
impl UnsafeUnpin for AdContext
impl !UnwindSafe for AdContext
Blanket Implementations§
Source§impl<T> BorrowMut<T> for Twhere
T: ?Sized,
impl<T> BorrowMut<T> for Twhere
T: ?Sized,
Source§fn borrow_mut(&mut self) -> &mut T
fn borrow_mut(&mut self) -> &mut T
Source§impl<T> CloneToUninit for Twhere
T: Clone,
impl<T> CloneToUninit for Twhere
T: Clone,
Source§impl<T> IntoEither for T
impl<T> IntoEither for T
Source§fn into_either(self, into_left: bool) -> Either<Self, Self>
fn into_either(self, into_left: bool) -> Either<Self, Self>
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 moreSource§fn into_either_with<F>(self, into_left: F) -> Either<Self, Self>
fn into_either_with<F>(self, into_left: F) -> Either<Self, Self>
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