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.
- AdContext
Builder - Builder for
AdContext. - AdContext
Cache Stats - Stats for caches owned by an
AdContext. - AdTransform
Cache Limits - Retention limits for AD transform graph caches.
- Context
Id - Opaque identifier for an eager AD runtime, used in
Error::ContextMismatch. - CpuPlacement
Bound Eager - Placement-selected CPU view of one
EagerRuntime. - DotGeneral
Config - DotGeneral dimension configuration.
- Eager
NoGrad Guard - Scope guard that temporarily disables eager operation recording.
- Eager
Runtime - Shared eager execution context for tensors on a backend.
- Eager
Runtime Cache Stats - Stats for caches owned by an
EagerRuntime. - Eager
Session - An eager runtime and its borrowed backend session for one execution boundary.
- Eager
Slice Builder - Rank-preserving eager tensor slicing builder.
- Eager
Tensor - Eager tensor with reverse-mode autodiff over concrete tensor values.
- Eager
Trace Capture Guard - Scope guard that keeps semantic-trace recording active for untracked intermediates.
- Gather
Config - StableHLO gather dimension configuration.
- Gradient
Value - Read-only retained gradient value.
- Gradients
- Move-only accumulated gradient bundle backed by one allocation group.
- PadConfig
- StableHLO pad configuration.
- Scatter
Config - StableHLO scatter dimension configuration.
- Slice
Config - Slice configuration.
- Tensor
- Dynamic tensor over the supported scalar types.
- Value
Guard - A read-only value view retained by an eager tensor record.
Enums§
- Compare
Dir - Comparison direction.
- DType
- Runtime scalar dtype tag.
- Error
- Errors produced by einsum, eval, and other tenferro operations.
- Error
Phase - Phase at which a runtime failure was discovered.
- Into
Value Error - Error returned when a value cannot be consumed without changing its owner.
Type Aliases§
- Result
- Result type alias for tenferro operations.