tidu/lib.rs
1//! Automatic-differentiation transforms for primitive computation graphs.
2//!
3//! `tidu` is for downstream crates that define primitive operations, local AD
4//! rules, graph runtimes, or eager tensor frontends. It does not define tensor
5//! operations itself. Instead, downstream primitive sets implement [`Primitive`],
6//! then call the graph transforms here to build new primitive computation
7//! graphs.
8//!
9//! The main transforms are:
10//!
11//! - [`linearize`], which builds a graph for a Jacobian-vector product (JVP)
12//! of selected outputs with respect to selected inputs.
13//! - [`linear_transpose`], which transposes a linearized graph so cotangents
14//! can flow backward through the corresponding linear map.
15//! - [`eager::backward`], which supports downstream eager frontends that
16//! record graph invocations and want a reverse-mode `backward()` workflow.
17//!
18//! These transforms propagate [`ADRuleError`] for missing primitive or
19//! extension AD rules.
20//!
21//! See the repository `docs/` tree for the terminology guide, tutorials, and
22//! implementer guides.
23//!
24//! # Examples
25//!
26//! ```ignore
27//! use computegraph::resolve::resolve;
28//! use tidu::{linear_transpose, linearize};
29//!
30//! let view = resolve(vec![source_graph]);
31//! let mut ctx = ();
32//! let aliases = std::collections::HashMap::new();
33//! let linear = linearize(&view, &[output_key], &[input_key], 1, &mut ctx, &aliases)?;
34//! let _transposed = linear_transpose(&linear, &mut ctx)?;
35//! # Ok::<(), tidu::ADRuleError>(())
36//! ```
37
38pub mod eager;
39mod linear_transpose;
40mod linearize;
41mod linearized_graph;
42mod primitive_graph;
43pub mod rules;
44
45pub use linear_transpose::{linear_transpose, linear_transpose_with_builder};
46pub use linearize::linearize;
47pub use linearized_graph::LinearizedGraph;
48pub use primitive_graph::PrimitiveGraph;
49pub use rules::{
50 ADKey, ADRuleError, ADRuleKind, ADRuleResult, DiffPassId, Primitive, PrimitiveBuilder,
51 PrimitiveValue,
52};