Skip to main content

EagerSessionEinsumExt

Trait EagerSessionEinsumExt 

Source
pub trait EagerSessionEinsumExt {
    // Required methods
    fn einsum(
        &mut self,
        inputs: &[&EagerTensor],
        subscripts: &str,
    ) -> Result<EagerTensor>;
    fn einsum_notation(
        &mut self,
        inputs: &[&EagerTensor],
        notation: &EinsumNotation,
    ) -> Result<EagerTensor>;
    fn einsum_subscripts(
        &mut self,
        inputs: &[&EagerTensor],
        subscripts: &EinsumSubscripts,
    ) -> Result<EagerTensor>;
    fn tensordot(
        &mut self,
        lhs: &EagerTensor,
        rhs: &EagerTensor,
        axes: TensorDotAxes<'_>,
    ) -> Result<EagerTensor>;
}
Expand description

Eager einsum and tensordot on a runtime-bound borrowed eager session.

Every operation runs inside the caller’s session, so it composes with other eager operations in the same tenferro_ad::EagerRuntime::with_eager_session callback and never reopens the runtime. The calling thread’s no_grad and capture_trace modes govern it like any other eager operation.

§Examples

use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
use tenferro_einsum::EagerSessionEinsumExt;

let ctx = EagerRuntime::new()?;
let a = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6])?, ctx.clone())?;
let b = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12])?, ctx.clone())?;
let product = ctx.with_eager_session(|session| {
    let c = session.einsum(&[&a, &b], "ij,jk->ik")?;
    session.einsum(&[&c], "ij->")
})?;
assert_eq!(product.value()?.as_slice::<f64>()?, &[24.0]);

Required Methods§

Source

fn einsum( &mut self, inputs: &[&EagerTensor], subscripts: &str, ) -> Result<EagerTensor>

Execute an einsum from string notation.

§Examples
use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
use tenferro_einsum::EagerSessionEinsumExt;

let ctx = EagerRuntime::new()?;
let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0])?, ctx.clone())?;
let dot = ctx.with_eager_session(|session| session.einsum(&[&x, &x], "i,i->"))?;
assert_eq!(dot.value()?.as_slice::<f64>()?, &[13.0]);
§Errors

Returns Error::InvalidSubscripts for malformed notation, Error::Validation for rank/shape/dtype mismatches, or Error::Planning / Error::Runtime for contraction planning and execution failures, including inputs owned by another runtime.

Source

fn einsum_notation( &mut self, inputs: &[&EagerTensor], notation: &EinsumNotation, ) -> Result<EagerTensor>

Execute an einsum from rank-unresolved notation.

§Examples
use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
use tenferro_einsum::{EagerSessionEinsumExt, EinsumAxis, EinsumNotation};

let ctx = EagerRuntime::new()?;
let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0])?, ctx.clone())?;
// `...->...` keeps every axis the ellipsis covers.
let notation = EinsumNotation::new(&[&[EinsumAxis::Ellipsis]], &[EinsumAxis::Ellipsis]);
let same = ctx.with_eager_session(|session| session.einsum_notation(&[&x], &notation))?;
assert_eq!(same.value()?.as_slice::<f64>()?, &[2.0, 3.0]);
§Errors

Returns a typed validation, planning, or runtime error when notation or execution is invalid.

Source

fn einsum_subscripts( &mut self, inputs: &[&EagerTensor], subscripts: &EinsumSubscripts, ) -> Result<EagerTensor>

Execute an einsum from parsed integer labels.

§Examples
use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
use tenferro_einsum::{EagerSessionEinsumExt, EinsumSubscripts};

let ctx = EagerRuntime::new()?;
let x = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2], vec![2.0_f64, 3.0])?, ctx.clone())?;
let subscripts = EinsumSubscripts::new(&[&[0], &[0]], &[]);
let dot = ctx.with_eager_session(|session| session.einsum_subscripts(&[&x, &x], &subscripts))?;
assert_eq!(dot.value()?.as_slice::<f64>()?, &[13.0]);
§Errors

Returns Error::Validation for rank/shape/dtype mismatches, Error::Planning for an invalid contraction plan, or Error::Runtime for extension registration or backend execution failures.

Source

fn tensordot( &mut self, lhs: &EagerTensor, rhs: &EagerTensor, axes: TensorDotAxes<'_>, ) -> Result<EagerTensor>

Contract two eager tensors over the requested axes.

§Examples
use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
use tenferro_einsum::{EagerSessionEinsumExt, TensorDotAxes};

let ctx = EagerRuntime::new()?;
let a = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![2, 3], vec![1.0_f64; 6])?, ctx.clone())?;
let b = EagerTensor::from_tensor_in(Tensor::from_vec_col_major(vec![3, 4], vec![1.0_f64; 12])?, ctx.clone())?;
let c = ctx.with_eager_session(|session| session.tensordot(&a, &b, TensorDotAxes::Count(1)))?;
assert_eq!(c.shape(), &[2, 4]);
§Errors

Returns Error::Validation for invalid axes or mismatched contracted extents, or Error::Runtime for execution failures.

Dyn Compatibility§

This trait is dyn compatible.

In older versions of Rust, dyn compatibility was called "object safety".

Implementations on Foreign Types§

Source§

impl EagerSessionEinsumExt for EagerSession<'_>

Source§

fn einsum( &mut self, inputs: &[&EagerTensor], subscripts: &str, ) -> Result<EagerTensor>

Source§

fn einsum_notation( &mut self, inputs: &[&EagerTensor], notation: &EinsumNotation, ) -> Result<EagerTensor>

Source§

fn einsum_subscripts( &mut self, inputs: &[&EagerTensor], subscripts: &EinsumSubscripts, ) -> Result<EagerTensor>

Source§

fn tensordot( &mut self, lhs: &EagerTensor, rhs: &EagerTensor, axes: TensorDotAxes<'_>, ) -> Result<EagerTensor>

Implementors§