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§
Sourcefn einsum(
&mut self,
inputs: &[&EagerTensor],
subscripts: &str,
) -> Result<EagerTensor>
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.
Sourcefn einsum_notation(
&mut self,
inputs: &[&EagerTensor],
notation: &EinsumNotation,
) -> Result<EagerTensor>
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], ¬ation))?;
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.
Sourcefn einsum_subscripts(
&mut self,
inputs: &[&EagerTensor],
subscripts: &EinsumSubscripts,
) -> Result<EagerTensor>
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.
Sourcefn tensordot(
&mut self,
lhs: &EagerTensor,
rhs: &EagerTensor,
axes: TensorDotAxes<'_>,
) -> Result<EagerTensor>
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".