Skip to main content

ContractionTree

Struct ContractionTree 

Source
pub struct ContractionTree { /* private fields */ }
Expand description

Contraction tree determining pairwise contraction order for N-ary einsum.

When contracting more than two tensors, the order in which pairwise contractions are performed significantly affects performance. ContractionTree encodes this order as a binary tree.

§Optimization

Use ContractionTree::optimize for automatic cost-based optimization (e.g., greedy algorithm based on tensor sizes), or ContractionTree::from_pairs for manual specification.

Implementations§

Source§

impl ContractionTree

Source

pub fn optimize(subscripts: &Subscripts, shapes: &[&[usize]]) -> Result<Self>

Automatically compute an optimized contraction order.

Uses a cost-based heuristic (greedy algorithm) to determine the pairwise contraction sequence that minimizes total operation count. The path is deterministic: for fixed subscripts and shapes it is the same in every process (see ContractionOptimizerOptions).

§Arguments
  • subscripts — Einsum subscripts for all tensors
  • shapes — Shape of each input tensor
§Examples
use tenferro_einsum::{ContractionTree, Subscripts};

let subs = Subscripts::parse("abcdef,bf,cf,df,ef->f").unwrap();
let shapes = [&[2, 3, 2, 4, 3, 9][..], &[3, 9], &[2, 9], &[4, 9], &[3, 9]];
let tree = ContractionTree::optimize(&subs, &shapes).unwrap();
assert_eq!(tree.step_count(), 4);
// Equal-cost candidates are broken by operand index, never by hashing.
let again = ContractionTree::optimize(&subs, &shapes).unwrap();
for step in 0..tree.step_count() {
    assert_eq!(tree.step_pair(step), again.step_pair(step));
}
§Errors

Returns Error::Validation when subscripts and shapes have different ranks or incompatible dimensions, or Error::Planning when no valid contraction order can be constructed.

Source

pub fn optimize_with_options( subscripts: &Subscripts, shapes: &[&[usize]], options: &ContractionOptimizerOptions, ) -> Result<Self>

Automatically compute an optimized contraction order with explicit planner options.

For three or more operands with an annealing schedule (niters > 0 and non-empty betas), this routes planning through TreeSA using the provided configuration. Without annealing, including the default options, TreeSA would return its greedy initializer unchanged, so the deterministic greedy planner runs directly and the path is identical in every process. One or two operands need no ordering search; their trees are built directly after validating the options.

§Errors

Returns Error::Validation for rank, shape, or dimension mismatches, or Error::Planning when planner options such as ntrials are invalid or no contraction order can be constructed.

Source

pub fn from_pairs( subscripts: &Subscripts, shapes: &[&[usize]], pairs: &[(usize, usize)], ) -> Result<Self>

Manually build a contraction tree from a pairwise contraction sequence.

Each pair (i, j) specifies which two tensors (or intermediate results) to contract next. Intermediate results are assigned indices starting from the number of input tensors.

§Arguments
  • subscripts — Einsum subscripts for all tensors
  • shapes — Shape of each input tensor
  • pairs — Ordered list of pairwise contractions
§Examples
use tenferro_einsum::{ContractionTree, Subscripts};

// Three tensors: A[ij] B[jk] C[kl] -> D[il]
// Contract B and C first, then A with the result:
let subs = Subscripts::new(&[&[0, 1], &[1, 2], &[2, 3]], &[0, 3]);
let shapes = [&[3, 4][..], &[4, 5], &[5, 6]];
let tree = ContractionTree::from_pairs(
    &subs,
    &shapes,
    &[(1, 2), (0, 3)],  // B*C -> T(index=3), then A*T -> D
).unwrap();
§Errors

Returns Error::Planning when the pair count, operand indices, or intermediate sequence is invalid, or Error::Validation when the supplied shapes do not match the subscripts.

Source

pub fn step_count(&self) -> usize

Return the number of pairwise contraction steps in this tree.

§Examples
use tenferro_einsum::{ContractionTree, Subscripts};

let subs = Subscripts::new(&[&[0, 1], &[1, 2], &[2, 3]], &[0, 3]);
let tree = ContractionTree::from_pairs(
    &subs,
    &[&[2, 2], &[2, 2], &[2, 2]],
    &[(1, 2), (0, 3)],
)
.unwrap();
assert_eq!(tree.step_count(), 2);
Source

pub fn step_pair(&self, step_idx: usize) -> Option<(usize, usize)>

Return the operand indices for a pairwise contraction step.

The returned indices refer to the original inputs (0..input_count) and then to intermediates (input_count..) produced by earlier steps.

§Examples
use tenferro_einsum::{ContractionTree, Subscripts};

let subs = Subscripts::new(&[&[0, 1], &[1, 2], &[2, 3]], &[0, 3]);
let tree = ContractionTree::from_pairs(
    &subs,
    &[&[2, 2], &[2, 2], &[2, 2]],
    &[(1, 2), (0, 3)],
)
.unwrap();
assert_eq!(tree.step_pair(0), Some((1, 2)));
Source

pub fn step_subscripts( &self, step_idx: usize, ) -> Option<(&[u32], &[u32], &[u32])>

Return the (lhs, rhs, output) subscripts for a pairwise step.

The output subscripts are the intermediate labels preserved after the contraction, or the final output labels on the last step.

§Examples
use tenferro_einsum::{ContractionTree, Subscripts};

let subs = Subscripts::new(&[&[0, 1], &[1, 2], &[2, 3]], &[0, 3]);
let tree = ContractionTree::from_pairs(
    &subs,
    &[&[2, 2], &[2, 2], &[2, 2]],
    &[(1, 2), (0, 3)],
)
.unwrap();
let (lhs, rhs, out) = tree.step_subscripts(0).unwrap();
assert_eq!(lhs, &[1, 2]);
assert_eq!(rhs, &[2, 3]);
assert_eq!(out, &[1, 3]);
Source

pub fn step_plan(&self, step_idx: usize) -> Option<PairwiseStepPlan<'_>>

Return the precomputed lowering plan for one pairwise contraction step.

§Examples
use tenferro_einsum::{ContractionTree, Subscripts};

let subs = Subscripts::new(&[&[0, 1], &[1, 2]], &[0, 2]);
let tree = ContractionTree::from_pairs(&subs, &[&[2, 3], &[3, 4]], &[(0, 1)]).unwrap();

assert_eq!(tree.step_plan(0).unwrap().gemm().m(), 2);

Trait Implementations§

Source§

impl Debug for ContractionTree

Source§

fn fmt(&self, f: &mut Formatter<'_>) -> Result

Formats the value using the given formatter. Read more

Auto Trait Implementations§

Blanket Implementations§

Source§

impl<T> Any for T
where T: 'static + ?Sized,

Source§

fn type_id(&self) -> TypeId

Gets the TypeId of self. Read more
Source§

impl<T> Borrow<T> for T
where T: ?Sized,

Source§

fn borrow(&self) -> &T

Immutably borrows from an owned value. Read more
Source§

impl<T> BorrowMut<T> for T
where T: ?Sized,

Source§

fn borrow_mut(&mut self) -> &mut T

Mutably borrows from an owned value. Read more
§

impl<T> ByRef<T> for T

§

fn by_ref(&self) -> &T

§

impl<ST, DT> CastableFrom<ST, Initialized, Initialized> for DT
where ST: ?Sized, DT: ?Sized,

§

impl<ST, DT> CastableFrom<ST, Uninit, Uninit> for DT
where ST: ?Sized, DT: ?Sized,

Source§

impl<T> From<T> for T

Source§

fn from(t: T) -> T

Returns the argument unchanged.

§

impl<T, U> Imply<T> for U
where T: ?Sized, U: ?Sized,

Source§

impl<T, U> Into<U> for T
where U: From<T>,

Source§

fn into(self) -> U

Calls U::from(self).

That is, this conversion is whatever the implementation of From<T> for U chooses to do.

Source§

impl<T> IntoEither for T

Source§

fn into_either(self, into_left: bool) -> Either<Self, Self> ⓘ

Converts self into a Left variant of Either<Self, Self> if into_left is true. Converts self into a Right variant of Either<Self, Self> otherwise. Read more
Source§

fn into_either_with<F>(self, into_left: F) -> Either<Self, Self> ⓘ
where F: FnOnce(&Self) -> bool,

Converts self into a Left variant of Either<Self, Self> if into_left(&self) returns true. Converts self into a Right variant of Either<Self, Self> otherwise. Read more
§

impl<T> MaybeSend for T
where T: Send,

§

impl<T> MaybeSendSync for T
where T: Send + Sync,

§

impl<T> MaybeSync for T
where T: Sync,

§

impl<T> Pointable for T

§

const ALIGN: usize

The alignment of pointer.
§

type Init = T

The type for initializers.
§

unsafe fn init(init: <T as Pointable>::Init) -> usize

Initializes a with the given initializer. Read more
§

unsafe fn deref<'a>(ptr: usize) -> &'a T

Dereferences the given pointer. Read more
§

unsafe fn deref_mut<'a>(ptr: usize) -> &'a mut T

Mutably dereferences the given pointer. Read more
§

unsafe fn drop(ptr: usize)

Drops the object pointed to by the given pointer. Read more
§

impl<T> Read<Exclusive, BecauseExclusive> for T
where T: ?Sized,

Source§

impl<T, U> TryFrom<U> for T
where U: Into<T>,

Source§

type Error = Infallible

The type returned in the event of a conversion error.
Source§

fn try_from(value: U) -> Result<T, <T as TryFrom<U>>::Error>

Performs the conversion.
Source§

impl<T, U> TryInto<U> for T
where U: TryFrom<T>,

Source§

type Error = <U as TryFrom<T>>::Error

The type returned in the event of a conversion error.
Source§

fn try_into(self) -> Result<U, <U as TryFrom<T>>::Error>

Performs the conversion.
§

impl<V, T> VZip<V> for T
where V: MultiLane<T>,

§

fn vzip(self) -> V