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//! # Cargo features
28//!
29//! | Feature | Enables |
30//! |---|---|
31//! | `cpu-faer` (default) and the other CPU provider features | Forwarded to `tenferro-cpu`; see its documentation. |
32//! | `autodiff` | The eager surface (`EagerSessionEinsumExt`) and AD rules. Adds the `tenferro-ad` dependency. |
33//! | `cuda` | CUDA execution through `tenferro-gpu`. |
34//! | `webgpu` | WebGPU/Metal execution through `tenferro-gpu` (a subset of operations). |
35//! | `rocm` | Placeholder; HIP/ROCm is not implemented. |
36//!
37//! For contractions without AD, use the concrete [`TensorEinsumExt`] /
38//! [`TypedTensorEinsumExt`] APIs or [`ConcreteEinsumPlan`] inside a backend
39//! session; they need no `autodiff`.
40//!
41//! # Examples
42//!
43//! ```
44//! use tenferro_einsum::{ContractionTree, Subscripts};
45//!
46//! let subs = Subscripts::parse("ij,jk->ik").unwrap();
47//! let tree = ContractionTree::optimize(&subs, &[&[2, 3], &[3, 4]]).unwrap();
48//! assert_eq!(tree.step_count(), 1);
49//! ```
50//!
51//! ```
52//! use tenferro_cpu::CpuBackend;
53//! use tenferro_einsum::TensorEinsumExt;
54//! use tenferro_tensor::{BackendSessionHost, Tensor};
55//!
56//! let a = Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6]).unwrap();
57//! let b = Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12]).unwrap();
58//! let mut backend = CpuBackend::new();
59//!
60//! let out = backend.with_backend_session(|session| {
61//! [&a, &b].einsum("ij,jk->ik", session)
62//! })??;
63//! assert_eq!(out.shape(), &[2, 4]);
64//! # Ok::<(), tenferro_einsum::Error>(())
65//! ```
66//!
67//! ```
68//! use tenferro_einsum::Subscripts;
69//!
70//! let trace = Subscripts::parse("ii->").unwrap();
71//! let diagonal = Subscripts::parse("ii->i").unwrap();
72//! let embedded = Subscripts::parse("i->ii").unwrap();
73//! let higher_rank = Subscripts::parse("iij->ij").unwrap();
74//!
75//! assert!(trace.output.is_empty());
76//! assert_eq!(diagonal.output, vec![b'i' as u32]);
77//! assert_eq!(embedded.output, vec![b'i' as u32, b'i' as u32]);
78//! assert_eq!(higher_rank.inputs[0], vec![b'i' as u32, b'i' as u32, b'j' as u32]);
79//! ```
80#![cfg_attr(docsrs, feature(doc_cfg))]
81
82mod binary_dot;
83mod builder;
84mod cache;
85mod concrete;
86mod eager;
87#[cfg(feature = "autodiff")]
88mod eager_ad;
89mod ellipsis;
90mod error;
91mod extension;
92pub mod lowering;
93mod optimize;
94mod planning;
95pub mod prelude;
96mod subscripts;
97mod syntax;
98mod tensordot;
99mod traced;
100#[cfg(test)]
101mod typed_eager;
102pub(crate) mod util;
103
104pub use cache::EINSUM_EXTENSION_FAMILY_ID;
105pub use concrete::{
106 ConcreteEinsumPlan, TensorEinsumExt, TensorEinsumIntoExt, TensorReadEinsumExt,
107 TensorReadEinsumIntoExt, TensorTensordotExt, TypedTensorEinsumExt, TypedTensorEinsumIntoExt,
108 TypedTensorReadEinsumExt, TypedTensorReadEinsumIntoExt, TypedTensorTensordotExt,
109};
110#[cfg(feature = "autodiff")]
111#[cfg_attr(docsrs, doc(cfg(feature = "autodiff")))]
112pub use eager_ad::EagerSessionEinsumExt;
113pub use error::{Error, PlanningError, Result};
114pub use extension::extension_module;
115#[cfg(feature = "autodiff")]
116pub use extension::semantic_ad_rules;
117pub use optimize::EinsumOptimize;
118pub use planning::tree::{ContractionOptimizerOptions, ContractionTree};
119pub use subscripts::{
120 parse_einsum_notation, parse_einsum_subscripts, EinsumAxis, EinsumNotation, EinsumSubscripts,
121};
122pub use syntax::nested::NestedEinsum;
123pub use syntax::subscripts::Subscripts;
124pub use tensordot::TensorDotAxes;
125pub use traced::{TraceContextEinsumExt, TracedTensorEinsumExt};
126
127#[cfg(test)]
128mod concrete_tests;
129#[cfg(test)]
130mod tests;
131#[cfg(test)]
132mod typed_eager_tests;