tenferro_einsum/lib.rs
1//! High-level einsum with N-ary contraction tree optimization.
2//!
3//! This crate provides:
4//!
5//! - **String notation**: `"ij,jk->ik"` (NumPy/PyTorch compatible)
6//! - **Parenthesized notation**: `"ij,(jk,kl)->il"` respects user-specified
7//! contraction order via [`NestedEinsum`]
8//! - **Integer label notation**: using `u32` labels
9//! - **Repeated labels**: `"ii->i"` extracts diagonals, `"ii->"` traces, and
10//! `"i->ii"` embeds a vector on a diagonal
11//! - **N-ary contraction**: Automatic or manual optimization of pairwise
12//! contraction order via [`ContractionTree`]
13//! - **Tensordot sugar**: NumPy-style axis-pair contraction extension methods,
14//! implemented as contraction sugar rather than as linear algebra APIs.
15//! - **Concrete execution**: backend-explicit [`TensorEinsumExt`],
16//! [`TypedTensorEinsumExt`], [`TensorReadEinsumExt`],
17//! [`TypedTensorReadEinsumExt`], and [`ConcreteEinsumPlan`] APIs for non-AD
18//! tensor values.
19//! - **Extension runtime**: traced einsum lowers to a registered tenferro
20//! extension runtime, keeping core op definitions small.
21//! - **Tensor extension traits**: graph-building and immediate-execution
22//! helpers are available as methods on `GraphCompiler`, concrete input
23//! slices/arrays, eager input slices/arrays, and tensor receivers.
24//!
25//! # Examples
26//!
27//! ```
28//! use tenferro_einsum::{ContractionTree, Subscripts};
29//!
30//! let subs = Subscripts::parse("ij,jk->ik").unwrap();
31//! let tree = ContractionTree::optimize(&subs, &[&[2, 3], &[3, 4]]).unwrap();
32//! assert_eq!(tree.step_count(), 1);
33//! ```
34//!
35//! ```
36//! use tenferro_cpu::CpuBackend;
37//! use tenferro_einsum::TensorEinsumExt;
38//! use tenferro_tensor::Tensor;
39//!
40//! let a = Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
41//! let b = Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12]).unwrap();
42//! let mut backend = CpuBackend::new();
43//!
44//! let out = [&a, &b].einsum("ij,jk->ik", &mut backend)?;
45//! assert_eq!(out.shape(), &[2, 4]);
46//! # Ok::<(), tenferro_einsum::Error>(())
47//! ```
48//!
49//! ```
50//! use tenferro_einsum::Subscripts;
51//!
52//! let trace = Subscripts::parse("ii->").unwrap();
53//! let diagonal = Subscripts::parse("ii->i").unwrap();
54//! let embedded = Subscripts::parse("i->ii").unwrap();
55//! let higher_rank = Subscripts::parse("iij->ij").unwrap();
56//!
57//! assert!(trace.output.is_empty());
58//! assert_eq!(diagonal.output, vec![b'i' as u32]);
59//! assert_eq!(embedded.output, vec![b'i' as u32, b'i' as u32]);
60//! assert_eq!(higher_rank.inputs[0], vec![b'i' as u32, b'i' as u32, b'j' as u32]);
61//! ```
62
63mod binary_dot;
64mod builder;
65mod cache;
66mod concrete;
67mod eager;
68#[cfg(feature = "autodiff")]
69mod eager_ad;
70mod error;
71mod extension;
72pub mod lowering;
73mod optimize;
74mod planning;
75mod subscripts;
76mod syntax;
77mod tensordot;
78mod traced;
79#[cfg(test)]
80mod typed_eager;
81pub(crate) mod util;
82
83pub use cache::EINSUM_EXTENSION_FAMILY_ID;
84pub use concrete::{
85 ConcreteEinsumPlan, TensorEinsumExt, TensorEinsumIntoExt, TensorReadEinsumExt,
86 TensorReadEinsumIntoExt, TensorTensordotExt, TypedTensorEinsumExt, TypedTensorEinsumIntoExt,
87 TypedTensorReadEinsumExt, TypedTensorReadEinsumIntoExt, TypedTensorTensordotExt,
88};
89#[cfg(feature = "autodiff")]
90pub use eager_ad::{EagerEinsumExt, EagerTensorEinsumExt};
91pub use error::{Error, PlanningError, Result};
92pub use extension::extension_module;
93#[cfg(feature = "autodiff")]
94pub use extension::semantic_ad_rules;
95pub use optimize::EinsumOptimize;
96pub use planning::tree::{ContractionOptimizerOptions, ContractionTree};
97pub use subscripts::{parse_einsum_subscripts, EinsumSubscripts};
98pub use syntax::nested::NestedEinsum;
99pub use syntax::subscripts::Subscripts;
100pub use tensordot::TensorDotAxes;
101pub use traced::{TraceContextEinsumExt, TracedTensorEinsumExt};
102
103#[cfg(test)]
104mod concrete_tests;
105#[cfg(test)]
106mod tests;
107#[cfg(test)]
108mod typed_eager_tests;