Skip to main content

Crate tenferro_ad

Crate tenferro_ad 

Source
Expand description

Automatic differentiation APIs for tenferro.

This crate is the explicit opt-in boundary for traced and eager automatic differentiation. Primal graph construction and execution live in tenferro-runtime; tensor storage lives in tenferro-tensor, and CPU execution lives in tenferro-cpu.

Use EagerRuntime and EagerTensor for PyTorch-style immediate execution where tracked variables accumulate gradients after backward(). Use TracedTensorAdExt or AdContext for JAX-style graph transforms such as grad, vjp, and jvp on tenferro_runtime::TracedTensor values. AdContext is the explicit place to add extension AD rule sets for operation-family crates such as tenferro-linalg.

User-facing guides live at https://tensor4all.org/tenferro-rs/guides/autodiff.html and https://tensor4all.org/tenferro-rs/guides/choosing-an-api.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 dx = ad.grad(&loss, &x).unwrap();
assert_eq!(dx.rank, 0);

§Errors

Error (the runtime error type, re-exported here) has public constructors for the failures downstream code reports itself: Error::invalid_argument, Error::unsupported, Error::dtype_mismatch, Error::validation, Error::runtime_state, and Error::extension for a typed source error. Each takes the operation name and an ErrorPhase. A tenferro_tensor::Error converts with From, so ? works on tensor-level results inside functions returning Result. Match on Error::kind rather than on variant shapes.

use tenferro_ad::{Error, ErrorPhase};

fn check_rank(rank: usize) -> tenferro_ad::Result<()> {
    if rank != 2 {
        return Err(Error::invalid_argument(
            "my_crate::attention",
            ErrorPhase::Execution,
            "query",
            format!("expected a rank-2 query, got rank {rank}"),
        ));
    }
    Ok(())
}

let error = check_rank(3).unwrap_err();
assert_eq!(
    error.kind(),
    tenferro_tensor::ErrorKind::Validation(tenferro_tensor::ValidationKind::InvalidArgument)
);
let unsupported = Error::unsupported("my_crate::op", ErrorPhase::Execution, "no GPU path yet");
assert_eq!(unsupported.kind(), tenferro_tensor::ErrorKind::Unsupported);
let tensor_error: Error = tenferro_tensor::Error::invalid_argument("op", "arg", "bad").into();
assert!(matches!(tensor_error.kind(), tenferro_tensor::ErrorKind::Validation(_)));

Re-exports§

pub use traced::TracedTensorAdExt;

Modules§

error
extension
Eager AD support for out-of-tree extension primitives.
prelude
Common eager and traced automatic-differentiation entry points.
semantic_extension
Semantic-program automatic-differentiation rules for extension operations.
semantic_transform
Whole-program automatic differentiation over semantic SSA programs.
traced

Structs§

AdContext
Explicit automatic-differentiation context.
AdContextBuilder
Builder for AdContext.
AdContextCacheStats
Stats for caches owned by an AdContext.
AdTransformCacheLimits
Retention limits for AD transform graph caches.
ContextId
Opaque identifier for an eager AD runtime, used in Error::ContextMismatch.
CpuPlacementBoundEager
Placement-selected CPU view of one EagerRuntime.
DotGeneralConfig
DotGeneral dimension configuration.
EagerNoGradGuard
Scope guard that temporarily disables eager operation recording.
EagerRuntime
Shared eager execution context for tensors on a backend.
EagerRuntimeCacheStats
Stats for caches owned by an EagerRuntime.
EagerSession
An eager runtime and its borrowed backend session for one execution boundary.
EagerSliceBuilder
Rank-preserving eager tensor slicing builder.
EagerTensor
Eager tensor with reverse-mode autodiff over concrete tensor values.
EagerTraceCaptureGuard
Scope guard that keeps semantic-trace recording active for untracked intermediates.
GatherConfig
StableHLO gather dimension configuration.
GradientValue
Read-only retained gradient value.
Gradients
Move-only accumulated gradient bundle backed by one allocation group.
PadConfig
StableHLO pad configuration.
ScatterConfig
StableHLO scatter dimension configuration.
SliceConfig
Slice configuration.
Tensor
Dynamic tensor over the supported scalar types.
ValueGuard
A read-only value view retained by an eager tensor record.

Enums§

CompareDir
Comparison direction.
DType
Runtime scalar dtype tag.
Error
Errors produced by einsum, eval, and other tenferro operations.
ErrorPhase
Phase at which a runtime failure was discovered.
IntoValueError
Error returned when a value cannot be consumed without changing its owner.

Type Aliases§

Result
Result type alias for tenferro operations.