Skip to main content

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