use tenferro_tensor::{
DType, DotGeneralAccumulation, Tensor, TensorBackend, TensorRead, TensorScalar, TensorWrite,
TypedTensor, TypedTensorView, TypedTensorWrite,
};
use crate::eager::{
eager_einsum_exec, eager_einsum_exec_read, eager_einsum_exec_read_into,
eager_einsum_exec_read_into_accum, eager_einsum_read_subscripts, eager_einsum_subscripts,
plan_subscripts,
};
use crate::TensorDotAxes;
use crate::{ContractionTree, EinsumSubscripts, Error, Result, Subscripts};
const TENSOR_EINSUM_OP: &str = "TensorEinsumExt::einsum";
const TENSOR_EINSUM_INTO_OP: &str = "TensorEinsumIntoExt::einsum_into";
const TENSOR_READ_EINSUM_OP: &str = "TensorReadEinsumExt::einsum_read";
const TENSOR_READ_EINSUM_INTO_OP: &str = "TensorReadEinsumIntoExt::einsum_read_into";
const TYPED_TENSOR_EINSUM_OP: &str = "TypedTensorEinsumExt::einsum";
const TYPED_TENSOR_EINSUM_INTO_OP: &str = "TypedTensorEinsumIntoExt::einsum_into";
const TYPED_TENSOR_READ_EINSUM_OP: &str = "TypedTensorReadEinsumExt::einsum_read";
const TYPED_TENSOR_READ_EINSUM_INTO_OP: &str = "TypedTensorReadEinsumIntoExt::einsum_read_into";
const PLAN_PREPARE_OP: &str = "ConcreteEinsumPlan::prepare";
const PLAN_EXECUTE_OP: &str = "ConcreteEinsumPlan::execute";
const TYPED_TENSOR_TENSORDOT_OP: &str = "TypedTensorTensordotExt::tensordot";
pub trait TensorTensordotExt {
fn tensordot<B: TensorBackend>(
&self,
rhs: &Tensor,
axes: TensorDotAxes<'_>,
backend: &mut B,
) -> Result<Tensor>;
}
impl TensorTensordotExt for Tensor {
fn tensordot<B: TensorBackend>(
&self,
rhs: &Tensor,
axes: TensorDotAxes<'_>,
backend: &mut B,
) -> Result<Tensor> {
let config =
crate::tensordot::dot_general_config(axes, self.shape().len(), rhs.shape().len())?;
crate::tensordot::validate_concrete_contract_dims(self.shape(), rhs.shape(), &config)?;
backend.dot_general(self, rhs, &config).map_err(Error::from)
}
}
pub trait TypedTensorTensordotExt<T: TensorScalar> {
fn tensordot<B: TensorBackend>(
&self,
rhs: &TypedTensor<T>,
axes: TensorDotAxes<'_>,
backend: &mut B,
) -> Result<TypedTensor<T>>;
}
impl<T: TensorScalar> TypedTensorTensordotExt<T> for TypedTensor<T> {
fn tensordot<B: TensorBackend>(
&self,
rhs: &TypedTensor<T>,
axes: TensorDotAxes<'_>,
backend: &mut B,
) -> Result<TypedTensor<T>> {
let config =
crate::tensordot::dot_general_config(axes, self.shape().len(), rhs.shape().len())?;
crate::tensordot::validate_concrete_contract_dims(self.shape(), rhs.shape(), &config)?;
let result = backend
.dot_general_read(T::tensor_read(self), T::tensor_read(rhs), &config)
.map_err(Error::from)?;
into_typed_result(result, TYPED_TENSOR_TENSORDOT_OP)
}
}
pub trait TensorEinsumExt {
fn einsum<B: TensorBackend>(&self, subscripts: &str, backend: &mut B) -> Result<Tensor>;
fn einsum_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<Tensor>;
}
impl TensorEinsumExt for [&Tensor] {
fn einsum<B: TensorBackend>(&self, subscripts: &str, backend: &mut B) -> Result<Tensor> {
let subscripts = parse_subscripts(subscripts, TENSOR_EINSUM_OP)?;
eager_einsum_subscripts(backend, self, &subscripts).map_err(Error::from)
}
fn einsum_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<Tensor> {
let subscripts = Subscripts::from(subscripts);
eager_einsum_subscripts(backend, self, &subscripts).map_err(Error::from)
}
}
impl<const N: usize> TensorEinsumExt for [&Tensor; N] {
fn einsum<B: TensorBackend>(&self, subscripts: &str, backend: &mut B) -> Result<Tensor> {
self.as_slice().einsum(subscripts, backend)
}
fn einsum_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<Tensor> {
self.as_slice().einsum_subscripts(subscripts, backend)
}
}
pub trait TensorEinsumIntoExt {
fn einsum_into<B: TensorBackend>(
&self,
subscripts: &str,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()>;
fn einsum_into_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()>;
}
impl TensorEinsumIntoExt for [&Tensor] {
fn einsum_into<B: TensorBackend>(
&self,
subscripts: &str,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()> {
let subscripts = parse_subscripts(subscripts, TENSOR_EINSUM_INTO_OP)?;
tensor_einsum_into_subscripts(backend, self, &subscripts, out, TENSOR_EINSUM_INTO_OP)
}
fn einsum_into_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()> {
let subscripts = Subscripts::from(subscripts);
tensor_einsum_into_subscripts(backend, self, &subscripts, out, TENSOR_EINSUM_INTO_OP)
}
}
impl<const N: usize> TensorEinsumIntoExt for [&Tensor; N] {
fn einsum_into<B: TensorBackend>(
&self,
subscripts: &str,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()> {
self.as_slice().einsum_into(subscripts, backend, out)
}
fn einsum_into_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()> {
self.as_slice()
.einsum_into_subscripts(subscripts, backend, out)
}
}
pub trait TypedTensorEinsumExt<T: TensorScalar> {
fn einsum<B: TensorBackend>(&self, subscripts: &str, backend: &mut B)
-> Result<TypedTensor<T>>;
fn einsum_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<TypedTensor<T>>;
}
impl<T: TensorScalar> TypedTensorEinsumExt<T> for [&TypedTensor<T>] {
fn einsum<B: TensorBackend>(
&self,
subscripts: &str,
backend: &mut B,
) -> Result<TypedTensor<T>> {
let subscripts = parse_subscripts(subscripts, TYPED_TENSOR_EINSUM_OP)?;
typed_einsum_subscripts(backend, self, &subscripts, TYPED_TENSOR_EINSUM_OP)
}
fn einsum_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<TypedTensor<T>> {
let subscripts = Subscripts::from(subscripts);
typed_einsum_subscripts(backend, self, &subscripts, TYPED_TENSOR_EINSUM_OP)
}
}
impl<T: TensorScalar, const N: usize> TypedTensorEinsumExt<T> for [&TypedTensor<T>; N] {
fn einsum<B: TensorBackend>(
&self,
subscripts: &str,
backend: &mut B,
) -> Result<TypedTensor<T>> {
self.as_slice().einsum(subscripts, backend)
}
fn einsum_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<TypedTensor<T>> {
self.as_slice().einsum_subscripts(subscripts, backend)
}
}
pub trait TypedTensorReadEinsumExt<T: TensorScalar> {
fn einsum_read<B: TensorBackend>(
&self,
subscripts: &str,
backend: &mut B,
) -> Result<TypedTensor<T>>;
fn einsum_read_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<TypedTensor<T>>;
}
impl<'a, T: TensorScalar> TypedTensorReadEinsumExt<T> for [TypedTensorView<'a, T>] {
fn einsum_read<B: TensorBackend>(
&self,
subscripts: &str,
backend: &mut B,
) -> Result<TypedTensor<T>> {
let subscripts = parse_subscripts(subscripts, TYPED_TENSOR_READ_EINSUM_OP)?;
typed_view_einsum_subscripts(backend, self, &subscripts, TYPED_TENSOR_READ_EINSUM_OP)
}
fn einsum_read_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<TypedTensor<T>> {
let subscripts = Subscripts::from(subscripts);
typed_view_einsum_subscripts(backend, self, &subscripts, TYPED_TENSOR_READ_EINSUM_OP)
}
}
impl<'a, T: TensorScalar, const N: usize> TypedTensorReadEinsumExt<T>
for [TypedTensorView<'a, T>; N]
{
fn einsum_read<B: TensorBackend>(
&self,
subscripts: &str,
backend: &mut B,
) -> Result<TypedTensor<T>> {
self.as_slice().einsum_read(subscripts, backend)
}
fn einsum_read_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<TypedTensor<T>> {
self.as_slice().einsum_read_subscripts(subscripts, backend)
}
}
pub trait TypedTensorEinsumIntoExt<T: TensorScalar> {
fn einsum_into<'out, B, O>(&self, subscripts: &str, backend: &mut B, out: O) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>;
fn einsum_into_subscripts<'out, B, O>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: O,
) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>;
}
impl<T: TensorScalar> TypedTensorEinsumIntoExt<T> for [&TypedTensor<T>] {
fn einsum_into<'out, B, O>(&self, subscripts: &str, backend: &mut B, out: O) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>,
{
let subscripts = parse_subscripts(subscripts, TYPED_TENSOR_EINSUM_INTO_OP)?;
typed_einsum_into_subscripts(
backend,
self,
&subscripts,
out.into(),
TYPED_TENSOR_EINSUM_INTO_OP,
)
}
fn einsum_into_subscripts<'out, B, O>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: O,
) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>,
{
let subscripts = Subscripts::from(subscripts);
typed_einsum_into_subscripts(
backend,
self,
&subscripts,
out.into(),
TYPED_TENSOR_EINSUM_INTO_OP,
)
}
}
impl<T: TensorScalar, const N: usize> TypedTensorEinsumIntoExt<T> for [&TypedTensor<T>; N] {
fn einsum_into<'out, B, O>(&self, subscripts: &str, backend: &mut B, out: O) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>,
{
self.as_slice().einsum_into(subscripts, backend, out)
}
fn einsum_into_subscripts<'out, B, O>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: O,
) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>,
{
self.as_slice()
.einsum_into_subscripts(subscripts, backend, out)
}
}
pub trait TypedTensorReadEinsumIntoExt<T: TensorScalar> {
fn einsum_read_into<'out, B, O>(&self, subscripts: &str, backend: &mut B, out: O) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>;
fn einsum_read_into_subscripts<'out, B, O>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: O,
) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>;
}
impl<'a, T: TensorScalar> TypedTensorReadEinsumIntoExt<T> for [TypedTensorView<'a, T>] {
fn einsum_read_into<'out, B, O>(&self, subscripts: &str, backend: &mut B, out: O) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>,
{
let subscripts = parse_subscripts(subscripts, TYPED_TENSOR_READ_EINSUM_INTO_OP)?;
typed_view_einsum_into_subscripts(
backend,
self,
&subscripts,
out.into(),
TYPED_TENSOR_READ_EINSUM_INTO_OP,
)
}
fn einsum_read_into_subscripts<'out, B, O>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: O,
) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>,
{
let subscripts = Subscripts::from(subscripts);
typed_view_einsum_into_subscripts(
backend,
self,
&subscripts,
out.into(),
TYPED_TENSOR_READ_EINSUM_INTO_OP,
)
}
}
impl<'a, T: TensorScalar, const N: usize> TypedTensorReadEinsumIntoExt<T>
for [TypedTensorView<'a, T>; N]
{
fn einsum_read_into<'out, B, O>(&self, subscripts: &str, backend: &mut B, out: O) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>,
{
self.as_slice().einsum_read_into(subscripts, backend, out)
}
fn einsum_read_into_subscripts<'out, B, O>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: O,
) -> Result<()>
where
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>,
{
self.as_slice()
.einsum_read_into_subscripts(subscripts, backend, out)
}
}
pub trait TensorReadEinsumExt {
fn einsum_read<B: TensorBackend>(&self, subscripts: &str, backend: &mut B) -> Result<Tensor>;
fn einsum_read_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<Tensor>;
}
impl<'a> TensorReadEinsumExt for [TensorRead<'a>] {
fn einsum_read<B: TensorBackend>(&self, subscripts: &str, backend: &mut B) -> Result<Tensor> {
let subscripts = parse_subscripts(subscripts, TENSOR_READ_EINSUM_OP)?;
eager_einsum_read_subscripts(backend, self, &subscripts).map_err(Error::from)
}
fn einsum_read_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<Tensor> {
let subscripts = Subscripts::from(subscripts);
eager_einsum_read_subscripts(backend, self, &subscripts).map_err(Error::from)
}
}
impl<'a, const N: usize> TensorReadEinsumExt for [TensorRead<'a>; N] {
fn einsum_read<B: TensorBackend>(&self, subscripts: &str, backend: &mut B) -> Result<Tensor> {
self.as_slice().einsum_read(subscripts, backend)
}
fn einsum_read_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
) -> Result<Tensor> {
self.as_slice().einsum_read_subscripts(subscripts, backend)
}
}
pub trait TensorReadEinsumIntoExt {
fn einsum_read_into<B: TensorBackend>(
&self,
subscripts: &str,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()>;
fn einsum_read_into_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()>;
}
impl<'a> TensorReadEinsumIntoExt for [TensorRead<'a>] {
fn einsum_read_into<B: TensorBackend>(
&self,
subscripts: &str,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()> {
let subscripts = parse_subscripts(subscripts, TENSOR_READ_EINSUM_INTO_OP)?;
tensor_read_einsum_into_subscripts(
backend,
self,
&subscripts,
out,
TENSOR_READ_EINSUM_INTO_OP,
)
}
fn einsum_read_into_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()> {
let subscripts = Subscripts::from(subscripts);
tensor_read_einsum_into_subscripts(
backend,
self,
&subscripts,
out,
TENSOR_READ_EINSUM_INTO_OP,
)
}
}
impl<'a, const N: usize> TensorReadEinsumIntoExt for [TensorRead<'a>; N] {
fn einsum_read_into<B: TensorBackend>(
&self,
subscripts: &str,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()> {
self.as_slice().einsum_read_into(subscripts, backend, out)
}
fn einsum_read_into_subscripts<B: TensorBackend>(
&self,
subscripts: &EinsumSubscripts,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()> {
self.as_slice()
.einsum_read_into_subscripts(subscripts, backend, out)
}
}
#[derive(Debug)]
pub struct ConcreteEinsumPlan {
tree: ContractionTree,
inputs: Vec<ConcreteEinsumInputSpec>,
}
impl ConcreteEinsumPlan {
pub fn prepare<'a, I>(inputs: I, subscripts: &str) -> Result<Self>
where
I: AsRef<[&'a Tensor]>,
{
let subscripts = parse_subscripts(subscripts, PLAN_PREPARE_OP)?;
Self::prepare_subscripts_internal(input_specs(inputs.as_ref()), &subscripts)
}
pub fn prepare_subscripts<'a, I>(inputs: I, subscripts: &EinsumSubscripts) -> Result<Self>
where
I: AsRef<[&'a Tensor]>,
{
let subscripts = Subscripts::from(subscripts);
Self::prepare_subscripts_internal(input_specs(inputs.as_ref()), &subscripts)
}
pub fn prepare_typed<'a, T, I>(inputs: I, subscripts: &str) -> Result<Self>
where
T: TensorScalar,
I: AsRef<[&'a TypedTensor<T>]>,
{
let subscripts = parse_subscripts(subscripts, PLAN_PREPARE_OP)?;
Self::prepare_subscripts_internal(typed_input_specs(inputs.as_ref()), &subscripts)
}
pub fn prepare_typed_subscripts<'a, T, I>(
inputs: I,
subscripts: &EinsumSubscripts,
) -> Result<Self>
where
T: TensorScalar,
I: AsRef<[&'a TypedTensor<T>]>,
{
let subscripts = Subscripts::from(subscripts);
Self::prepare_subscripts_internal(typed_input_specs(inputs.as_ref()), &subscripts)
}
pub fn prepare_read<'a, I>(inputs: I, subscripts: &str) -> Result<Self>
where
I: AsRef<[TensorRead<'a>]>,
{
let subscripts = parse_subscripts(subscripts, PLAN_PREPARE_OP)?;
Self::prepare_subscripts_internal(read_input_specs(inputs.as_ref()), &subscripts)
}
pub fn prepare_read_subscripts<'a, I>(inputs: I, subscripts: &EinsumSubscripts) -> Result<Self>
where
I: AsRef<[TensorRead<'a>]>,
{
let subscripts = Subscripts::from(subscripts);
Self::prepare_subscripts_internal(read_input_specs(inputs.as_ref()), &subscripts)
}
pub fn execute<'a, I, B>(&self, inputs: I, backend: &mut B) -> Result<Tensor>
where
I: AsRef<[&'a Tensor]>,
B: TensorBackend,
{
let inputs = inputs.as_ref();
self.validate_inputs(&input_specs(inputs), PLAN_EXECUTE_OP)?;
backend
.with_backend_session(|exec| eager_einsum_exec(exec, inputs, &self.tree))
.map_err(Error::from)
}
pub fn execute_typed<'a, T, I, B>(&self, inputs: I, backend: &mut B) -> Result<TypedTensor<T>>
where
T: TensorScalar,
I: AsRef<[&'a TypedTensor<T>]>,
B: TensorBackend,
{
let inputs = inputs.as_ref();
self.validate_inputs(&typed_input_specs(inputs), PLAN_EXECUTE_OP)?;
let reads: Vec<_> = inputs.iter().map(|tensor| T::tensor_read(tensor)).collect();
let result = backend
.with_backend_session(|exec| eager_einsum_exec_read(exec, &reads, &self.tree))?;
into_typed_result(result, PLAN_EXECUTE_OP)
}
pub fn execute_read<'a, I, B>(&self, inputs: I, backend: &mut B) -> Result<Tensor>
where
I: AsRef<[TensorRead<'a>]>,
B: TensorBackend,
{
let inputs = inputs.as_ref();
self.validate_inputs(&read_input_specs(inputs), PLAN_EXECUTE_OP)?;
backend
.with_backend_session(|exec| eager_einsum_exec_read(exec, inputs, &self.tree))
.map_err(Error::from)
}
pub fn execute_into<'a, I, B>(
&self,
inputs: I,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()>
where
I: AsRef<[&'a Tensor]>,
B: TensorBackend,
{
let inputs = inputs.as_ref();
let specs = input_specs(inputs);
self.validate_inputs(&specs, PLAN_EXECUTE_OP)?;
validate_output(&self.inputs, &self.tree, &out, PLAN_EXECUTE_OP)?;
let reads: Vec<_> = inputs
.iter()
.map(|tensor| TensorRead::from_tensor(tensor))
.collect();
backend
.with_backend_session(|exec| eager_einsum_exec_read_into(exec, &reads, &self.tree, out))
.map_err(Error::from)
}
pub fn execute_typed_into<'a, 'out, T, I, B, O>(
&self,
inputs: I,
backend: &mut B,
out: O,
) -> Result<()>
where
T: TensorScalar,
I: AsRef<[&'a TypedTensor<T>]>,
B: TensorBackend,
O: Into<TypedTensorWrite<'out, T>>,
{
let inputs = inputs.as_ref();
let specs = typed_input_specs(inputs);
self.validate_inputs(&specs, PLAN_EXECUTE_OP)?;
let out = out.into().into_tensor_write();
validate_output(&self.inputs, &self.tree, &out, PLAN_EXECUTE_OP)?;
let reads: Vec<_> = inputs.iter().map(|tensor| T::tensor_read(tensor)).collect();
backend
.with_backend_session(|exec| eager_einsum_exec_read_into(exec, &reads, &self.tree, out))
.map_err(Error::from)
}
pub fn execute_read_into<'a, I, B>(
&self,
inputs: I,
backend: &mut B,
out: TensorWrite<'_>,
) -> Result<()>
where
I: AsRef<[TensorRead<'a>]>,
B: TensorBackend,
{
let inputs = inputs.as_ref();
let specs = read_input_specs(inputs);
self.validate_inputs(&specs, PLAN_EXECUTE_OP)?;
validate_output(&self.inputs, &self.tree, &out, PLAN_EXECUTE_OP)?;
backend
.with_backend_session(|exec| eager_einsum_exec_read_into(exec, inputs, &self.tree, out))
.map_err(Error::from)
}
pub fn execute_read_into_accum<'a, I, B>(
&self,
inputs: I,
backend: &mut B,
accumulation: DotGeneralAccumulation,
out: TensorWrite<'_>,
) -> Result<()>
where
I: AsRef<[TensorRead<'a>]>,
B: TensorBackend,
{
let inputs = inputs.as_ref();
let specs = read_input_specs(inputs);
self.validate_inputs(&specs, PLAN_EXECUTE_OP)?;
validate_output(&self.inputs, &self.tree, &out, PLAN_EXECUTE_OP)?;
backend
.with_backend_session(|exec| {
eager_einsum_exec_read_into_accum(exec, inputs, &self.tree, accumulation, out)
})
.map_err(Error::from)
}
fn prepare_subscripts_internal(
inputs: Vec<ConcreteEinsumInputSpec>,
subscripts: &Subscripts,
) -> Result<Self> {
let shapes: Vec<&[usize]> = inputs.iter().map(|input| input.shape.as_slice()).collect();
let tree = plan_subscripts(subscripts, &shapes)?;
Ok(Self { tree, inputs })
}
fn validate_inputs(&self, actual: &[ConcreteEinsumInputSpec], op: &'static str) -> Result<()> {
if actual.len() != self.inputs.len() {
return Err(Error::invalid_argument(
op,
"inputs",
format!(
"prepared einsum expects {} inputs, got {}",
self.inputs.len(),
actual.len()
),
));
}
for (expected, actual) in self.inputs.iter().zip(actual.iter()) {
if expected.dtype != actual.dtype {
return Err(Error::dtype_mismatch(op, expected.dtype, actual.dtype));
}
if expected.shape != actual.shape {
return Err(Error::shape_mismatch(
op,
expected.shape.clone(),
actual.shape.clone(),
));
}
}
Ok(())
}
}
#[derive(Clone, Debug)]
struct ConcreteEinsumInputSpec {
dtype: DType,
shape: Vec<usize>,
}
fn parse_subscripts(subscripts: &str, _op: &'static str) -> Result<Subscripts> {
Subscripts::parse(subscripts)
}
fn input_specs(inputs: &[&Tensor]) -> Vec<ConcreteEinsumInputSpec> {
inputs
.iter()
.map(|tensor| ConcreteEinsumInputSpec {
dtype: tensor.dtype(),
shape: tensor.shape().to_vec(),
})
.collect()
}
fn typed_input_specs<T: TensorScalar>(inputs: &[&TypedTensor<T>]) -> Vec<ConcreteEinsumInputSpec> {
inputs
.iter()
.map(|tensor| ConcreteEinsumInputSpec {
dtype: T::dtype(),
shape: tensor.shape().to_vec(),
})
.collect()
}
fn typed_view_input_specs<T: TensorScalar>(
inputs: &[TypedTensorView<'_, T>],
) -> Vec<ConcreteEinsumInputSpec> {
inputs
.iter()
.map(|tensor| ConcreteEinsumInputSpec {
dtype: T::dtype(),
shape: tensor.shape().to_vec(),
})
.collect()
}
fn read_input_specs(inputs: &[TensorRead<'_>]) -> Vec<ConcreteEinsumInputSpec> {
inputs
.iter()
.map(|tensor| ConcreteEinsumInputSpec {
dtype: tensor.dtype(),
shape: tensor.shape().to_vec(),
})
.collect()
}
fn typed_view_einsum_subscripts<T: TensorScalar>(
backend: &mut impl TensorBackend,
inputs: &[TypedTensorView<'_, T>],
subscripts: &Subscripts,
op: &'static str,
) -> Result<TypedTensor<T>> {
let reads: Vec<_> = inputs
.iter()
.cloned()
.map(|view| TensorRead::from_view(T::tensor_view(view)))
.collect();
let result = eager_einsum_read_subscripts(backend, &reads, subscripts)?;
into_typed_result(result, op)
}
fn tensor_einsum_into_subscripts(
backend: &mut impl TensorBackend,
inputs: &[&Tensor],
subscripts: &Subscripts,
out: TensorWrite<'_>,
op: &'static str,
) -> Result<()> {
let specs = input_specs(inputs);
let plan = ConcreteEinsumPlan::prepare_subscripts_internal(specs.clone(), subscripts)?;
validate_output(&specs, &plan.tree, &out, op)?;
let reads: Vec<_> = inputs
.iter()
.map(|tensor| TensorRead::from_tensor(tensor))
.collect();
backend
.with_backend_session(|exec| eager_einsum_exec_read_into(exec, &reads, &plan.tree, out))
.map_err(Error::from)
}
fn typed_view_einsum_into_subscripts<T: TensorScalar>(
backend: &mut impl TensorBackend,
inputs: &[TypedTensorView<'_, T>],
subscripts: &Subscripts,
out: TypedTensorWrite<'_, T>,
op: &'static str,
) -> Result<()> {
let specs = typed_view_input_specs(inputs);
let plan = ConcreteEinsumPlan::prepare_subscripts_internal(specs.clone(), subscripts)?;
let out = out.into_tensor_write();
validate_output(&specs, &plan.tree, &out, op)?;
let reads: Vec<_> = inputs
.iter()
.cloned()
.map(|view| TensorRead::from_view(T::tensor_view(view)))
.collect();
backend
.with_backend_session(|exec| eager_einsum_exec_read_into(exec, &reads, &plan.tree, out))
.map_err(Error::from)
}
fn typed_einsum_into_subscripts<T: TensorScalar>(
backend: &mut impl TensorBackend,
inputs: &[&TypedTensor<T>],
subscripts: &Subscripts,
out: TypedTensorWrite<'_, T>,
op: &'static str,
) -> Result<()> {
let specs = typed_input_specs(inputs);
let plan = ConcreteEinsumPlan::prepare_subscripts_internal(specs.clone(), subscripts)?;
let out = out.into_tensor_write();
validate_output(&specs, &plan.tree, &out, op)?;
let reads: Vec<_> = inputs.iter().map(|tensor| T::tensor_read(tensor)).collect();
backend
.with_backend_session(|exec| eager_einsum_exec_read_into(exec, &reads, &plan.tree, out))
.map_err(Error::from)
}
fn tensor_read_einsum_into_subscripts(
backend: &mut impl TensorBackend,
inputs: &[TensorRead<'_>],
subscripts: &Subscripts,
out: TensorWrite<'_>,
op: &'static str,
) -> Result<()> {
let specs = read_input_specs(inputs);
let plan = ConcreteEinsumPlan::prepare_subscripts_internal(specs.clone(), subscripts)?;
validate_output(&specs, &plan.tree, &out, op)?;
backend
.with_backend_session(|exec| eager_einsum_exec_read_into(exec, inputs, &plan.tree, out))
.map_err(Error::from)
}
fn validate_output(
inputs: &[ConcreteEinsumInputSpec],
tree: &ContractionTree,
out: &TensorWrite<'_>,
op: &'static str,
) -> Result<()> {
let expected = output_spec(inputs, tree, op)?;
if out.dtype() != expected.dtype {
return Err(Error::dtype_mismatch(op, expected.dtype, out.dtype()));
}
if out.shape() != expected.shape.as_slice() {
return Err(Error::shape_mismatch(
op,
out.shape().to_vec(),
expected.shape.clone(),
));
}
Ok(())
}
fn output_spec(
inputs: &[ConcreteEinsumInputSpec],
tree: &ContractionTree,
op: &'static str,
) -> Result<ConcreteEinsumInputSpec> {
let dtype = inputs
.first()
.ok_or_else(|| {
Error::invalid_argument(op, "inputs", "einsum requires at least one input tensor")
})?
.dtype;
for input in inputs {
if input.dtype != dtype {
return Err(Error::dtype_mismatch(op, dtype, input.dtype));
}
}
let mut output_shape = Vec::with_capacity(tree.subscripts.output.len());
for &label in &tree.subscripts.output {
let mut found = None;
for (input, labels) in inputs.iter().zip(tree.subscripts.inputs.iter()) {
if labels.len() != input.shape.len() {
return Err(Error::rank_mismatch(op, labels.len(), input.shape.len()));
}
if let Some(axis) = labels.iter().position(|candidate| *candidate == label) {
found = Some(input.shape[axis]);
break;
}
}
let Some(extent) = found else {
return Err(Error::invalid_argument(
op,
"output labels",
format!("output label {label} is missing from inputs"),
));
};
output_shape.push(extent);
}
Ok(ConcreteEinsumInputSpec {
dtype,
shape: output_shape,
})
}
fn typed_einsum_subscripts<T: TensorScalar>(
backend: &mut impl TensorBackend,
inputs: &[&TypedTensor<T>],
subscripts: &Subscripts,
op: &'static str,
) -> Result<TypedTensor<T>> {
let reads: Vec<_> = inputs.iter().map(|tensor| T::tensor_read(tensor)).collect();
let result = eager_einsum_read_subscripts(backend, &reads, subscripts)?;
into_typed_result(result, op)
}
pub(crate) fn into_typed_result<T: TensorScalar>(
result: Tensor,
op: &'static str,
) -> Result<TypedTensor<T>> {
let actual = result.dtype();
T::into_typed(result).map_err(|_| Error::dtype_mismatch(op, T::dtype(), actual))
}