tenferro_ad/lib.rs
1//! Automatic differentiation APIs for tenferro.
2//!
3//! This crate is the explicit opt-in boundary for traced and eager automatic
4//! differentiation. Primal graph construction and execution live in
5//! `tenferro-runtime`; tensor storage lives in `tenferro-tensor`, and CPU
6//! execution lives in `tenferro-cpu`.
7//!
8//! Use [`EagerRuntime`] and [`EagerTensor`] for PyTorch-style immediate
9//! execution where tracked variables accumulate gradients after `backward()`.
10//! Use [`TracedTensorAdExt`] or [`AdContext`] for JAX-style graph transforms
11//! such as `grad`, `vjp`, and `jvp` on [`tenferro_runtime::TracedTensor`]
12//! values. `AdContext` is the explicit place to add extension AD rule sets for
13//! operation-family crates such as `tenferro-linalg`.
14//!
15//! User-facing guides live at
16//! <https://tensor4all.org/tenferro-rs/guides/autodiff.html> and
17//! <https://tensor4all.org/tenferro-rs/guides/choosing-an-api.html>.
18//!
19//! # Examples
20//!
21//! ```rust
22//! use tenferro_ad::AdContext;
23//! use tenferro_runtime::TracedTensor;
24//!
25//! let ad = AdContext::builder().build().unwrap();
26//! let x = TracedTensor::from_vec_col_major(vec![], vec![3.0_f64]).unwrap();
27//! let loss = (&x * &x).unwrap();
28//! let dx = ad.grad(&loss, &x).unwrap();
29//! assert_eq!(dx.rank, 0);
30//! ```
31//!
32//! # Errors
33//!
34//! [`Error`] (the runtime error type, re-exported here) has public
35//! constructors for the failures downstream code reports itself:
36//! [`Error::invalid_argument`], [`Error::unsupported`],
37//! [`Error::dtype_mismatch`], [`Error::validation`], [`Error::runtime_state`],
38//! and [`Error::extension`] for a typed source error. Each takes the operation
39//! name and an [`ErrorPhase`]. A `tenferro_tensor::Error` converts with
40//! `From`, so `?` works on tensor-level results inside functions returning
41//! [`Result`]. Match on [`Error::kind`] rather than on variant shapes.
42//!
43//! ```rust
44//! use tenferro_ad::{Error, ErrorPhase};
45//!
46//! fn check_rank(rank: usize) -> tenferro_ad::Result<()> {
47//! if rank != 2 {
48//! return Err(Error::invalid_argument(
49//! "my_crate::attention",
50//! ErrorPhase::Execution,
51//! "query",
52//! format!("expected a rank-2 query, got rank {rank}"),
53//! ));
54//! }
55//! Ok(())
56//! }
57//!
58//! let error = check_rank(3).unwrap_err();
59//! assert_eq!(
60//! error.kind(),
61//! tenferro_tensor::ErrorKind::Validation(tenferro_tensor::ValidationKind::InvalidArgument)
62//! );
63//! let unsupported = Error::unsupported("my_crate::op", ErrorPhase::Execution, "no GPU path yet");
64//! assert_eq!(unsupported.kind(), tenferro_tensor::ErrorKind::Unsupported);
65//! let tensor_error: Error = tenferro_tensor::Error::invalid_argument("op", "arg", "bad").into();
66//! assert!(matches!(tensor_error.kind(), tenferro_tensor::ErrorKind::Validation(_)));
67//! ```
68
69mod context;
70mod eager;
71mod eager_backend;
72pub(crate) mod eager_exec;
73pub(crate) mod eager_ops;
74pub mod extension;
75pub mod prelude;
76// semantic_compat removed in Unification 7.
77pub mod semantic_extension;
78pub mod semantic_transform;
79mod shape_packing;
80pub mod traced;
81mod transform_cache;
82
83pub use context::{AdContext, AdContextBuilder, AdContextCacheStats};
84pub use eager::{
85 CpuPlacementBoundEager, EagerNoGradGuard, EagerRuntime, EagerRuntimeCacheStats, EagerSession,
86 EagerTensor, EagerTraceCaptureGuard, GradientValue, Gradients, IntoValueError, ValueGuard,
87};
88pub use shape_packing::EagerSliceBuilder;
89pub(crate) use tenferro_runtime::{extension_cache, scalar_semantics};
90pub use transform_cache::AdTransformCacheLimits;
91pub(crate) mod shape_infer {
92 pub use tenferro_runtime::extension::{
93 promote_dtype, promote_dtype_for_binary_op, promote_dtypes,
94 };
95}
96pub use tenferro_runtime::{
97 CompareDir, DType, DotGeneralConfig, GatherConfig, PadConfig, ScatterConfig, SliceConfig,
98 Tensor,
99};
100pub use traced::TracedTensorAdExt;
101
102pub use tenferro_runtime::{ContextId, Error, ErrorPhase, Result};
103
104pub mod error {
105 pub use tenferro_runtime::{ContextId, Error, ErrorPhase, Result};
106}
107
108pub(crate) mod metadata {
109 pub use tenferro_runtime::ad_support::tensor_meta_from_tensor;
110}