use std::collections::{HashMap, HashSet};
use std::time::{Duration, Instant};
use tenferro_tensor::{
BackendSession, DotGeneralAccumulation, DotGeneralConfig, Error, Result, Tensor, TensorBackend,
TensorRank, TensorRead, TensorView, TensorWrite, TypedTensorView,
};
use crate::binary_dot::{try_build_binary_dot_plan, BinaryDotPlan};
use crate::util::map_label_occurrences;
use crate::{ContractionTree, Subscripts};
mod profile;
use profile::{
eager_einsum_profile_enabled, eager_einsum_trace_enabled, maybe_print_eager_einsum_profile,
profile_eager_einsum_section, record_eager_einsum_profile,
};
const EAGER_EINSUM_OP: &str = "eager_einsum";
#[allow(clippy::large_enum_variant)]
enum TensorValue<'a> {
Borrowed(&'a Tensor),
View(TensorView<'a>),
Owned(Tensor),
}
impl TensorValue<'_> {
fn as_tensor(&self) -> Option<&Tensor> {
match self {
Self::Borrowed(tensor) => Some(tensor),
Self::View(_) => None,
Self::Owned(tensor) => Some(tensor),
}
}
fn into_tensor(self, exec: &mut dyn BackendSession) -> Result<Tensor> {
match self {
Self::Borrowed(tensor) => tensor.duplicate(),
Self::View(view) => exec.to_contiguous_read(TensorRead::from_view(view)),
Self::Owned(tensor) => Ok(tensor),
}
}
fn tensor_read(&self) -> TensorRead<'_> {
match self {
Self::Borrowed(tensor) => TensorRead::from_tensor(tensor),
Self::View(view) => TensorRead::from_view(view.clone()),
Self::Owned(tensor) => TensorRead::from_tensor(tensor),
}
}
fn tensor_view(&self) -> TensorView<'_> {
match self {
Self::Borrowed(tensor) => tensor_as_view(tensor),
Self::View(view) => view.clone(),
Self::Owned(tensor) => tensor_as_view(tensor),
}
}
fn reclaim_if_owned(self, exec: &mut dyn BackendSession) {
if let Self::Owned(tensor) = self {
exec.reclaim_buffer(tensor);
}
}
}
fn tensor_as_view(tensor: &Tensor) -> TensorView<'_> {
match tensor {
Tensor::F32(tensor) => TensorView::F32(tensor.as_view()),
Tensor::F64(tensor) => TensorView::F64(tensor.as_view()),
Tensor::I32(tensor) => TensorView::I32(tensor.as_view()),
Tensor::I64(tensor) => TensorView::I64(tensor.as_view()),
Tensor::Bool(tensor) => TensorView::Bool(tensor.as_view()),
Tensor::C32(tensor) => TensorView::C32(tensor.as_view()),
Tensor::C64(tensor) => TensorView::C64(tensor.as_view()),
}
}
fn tensor_view_has_backend_buffer(view: &TensorView<'_>) -> bool {
match view {
TensorView::F32(view) => view.backend_buffer().is_some(),
TensorView::F64(view) => view.backend_buffer().is_some(),
TensorView::I32(view) => view.backend_buffer().is_some(),
TensorView::I64(view) => view.backend_buffer().is_some(),
TensorView::Bool(view) => view.backend_buffer().is_some(),
TensorView::C32(view) => view.backend_buffer().is_some(),
TensorView::C64(view) => view.backend_buffer().is_some(),
}
}
fn broadcast_shape_strides<T: 'static, R: TensorRank>(
view: &TypedTensorView<'_, T, R>,
shape: &[usize],
dims: &[usize],
) -> Result<(Vec<usize>, Vec<isize>)> {
if dims.len() != view.shape().len() {
return Err(Error::rank_mismatch(
EAGER_EINSUM_OP,
view.shape().len(),
dims.len(),
));
}
let mut seen = vec![false; shape.len()];
let mut strides = vec![0isize; shape.len()];
for (src_axis, &dst_axis) in dims.iter().enumerate() {
if dst_axis >= shape.len() {
return Err(Error::axis_out_of_bounds(
EAGER_EINSUM_OP,
dst_axis,
shape.len(),
));
}
if seen[dst_axis] {
return Err(Error::duplicate_axis(
EAGER_EINSUM_OP,
dst_axis,
"broadcast dims",
));
}
seen[dst_axis] = true;
let source_dim = view.shape()[src_axis];
let target_dim = shape[dst_axis];
if source_dim != target_dim && source_dim != 1 {
return Err(Error::shape_mismatch(EAGER_EINSUM_OP, view.shape(), shape));
}
if source_dim == target_dim {
strides[dst_axis] = view.strides()[src_axis];
}
}
Ok((shape.to_vec(), strides))
}
fn broadcast_typed_view<'a, T: 'static, R: TensorRank>(
view: TypedTensorView<'a, T, R>,
shape: &[usize],
dims: &[usize],
) -> Result<TypedTensorView<'a, T>> {
let (shape, strides) = broadcast_shape_strides(&view, shape, dims)?;
TypedTensorView::from_slice(shape, strides, view.offset(), view.host_storage()?)
}
fn broadcast_tensor_view<'a>(
view: TensorView<'a>,
shape: &[usize],
dims: &[usize],
) -> Result<TensorView<'a>> {
match view {
TensorView::F32(view) => Ok(TensorView::F32(broadcast_typed_view(view, shape, dims)?)),
TensorView::F64(view) => Ok(TensorView::F64(broadcast_typed_view(view, shape, dims)?)),
TensorView::I32(view) => Ok(TensorView::I32(broadcast_typed_view(view, shape, dims)?)),
TensorView::I64(view) => Ok(TensorView::I64(broadcast_typed_view(view, shape, dims)?)),
TensorView::Bool(view) => Ok(TensorView::Bool(broadcast_typed_view(view, shape, dims)?)),
TensorView::C32(view) => Ok(TensorView::C32(broadcast_typed_view(view, shape, dims)?)),
TensorView::C64(view) => Ok(TensorView::C64(broadcast_typed_view(view, shape, dims)?)),
}
}
fn try_broadcast_tensor_read<'a>(
value: &'a TensorValue<'_>,
shape: &[usize],
dims: &[usize],
) -> Option<Result<TensorRead<'a>>> {
let view = value.tensor_view();
if tensor_view_has_backend_buffer(&view) {
return None;
}
Some(broadcast_tensor_view(view, shape, dims).map(TensorRead::from_view))
}
fn select_outer_product_label_order(
canonical_labels: &[u32],
target_labels: Option<&[u32]>,
) -> Vec<u32> {
let Some(target_labels) = target_labels else {
return canonical_labels.to_vec();
};
if target_labels.len() != canonical_labels.len() {
return canonical_labels.to_vec();
}
let mut used = vec![false; canonical_labels.len()];
for &label in target_labels {
let Some(axis) = canonical_labels
.iter()
.enumerate()
.find_map(|(axis, candidate)| (*candidate == label && !used[axis]).then_some(axis))
else {
return canonical_labels.to_vec();
};
used[axis] = true;
}
target_labels.to_vec()
}
struct LabeledTensor<'a> {
tensor: TensorValue<'a>,
labels: Vec<u32>,
}
impl LabeledTensor<'_> {
fn tensor(&self) -> Option<&Tensor> {
self.tensor.as_tensor()
}
fn tensor_read(&self) -> TensorRead<'_> {
self.tensor.tensor_read()
}
fn tensor_owned(&self, exec: &mut dyn BackendSession) -> Result<Tensor> {
match self.tensor() {
Some(tensor) => tensor.duplicate(),
None => exec.to_contiguous_read(self.tensor_read()),
}
}
fn shape(&self) -> &[usize] {
match &self.tensor {
TensorValue::Borrowed(tensor) => tensor.shape(),
TensorValue::View(view) => view.shape(),
TensorValue::Owned(tensor) => tensor.shape(),
}
}
fn reclaim_if_owned(self, exec: &mut dyn BackendSession) {
self.tensor.reclaim_if_owned(exec);
}
}
fn execute_binary_dot_fast_plan<'a>(
exec: &mut dyn BackendSession,
lhs: LabeledTensor<'a>,
rhs: LabeledTensor<'a>,
plan: BinaryDotPlan,
reorder_result: bool,
) -> Result<LabeledTensor<'a>> {
let tensor = profile_eager_einsum_section("binary.fast_dot_general", || {
exec.dot_general_read(lhs.tensor_read(), rhs.tensor_read(), &plan.config)
})?;
lhs.reclaim_if_owned(exec);
rhs.reclaim_if_owned(exec);
let result = LabeledTensor {
tensor: TensorValue::Owned(tensor),
labels: plan.result_labels,
};
if !reorder_result || result.labels == plan.target_labels {
return Ok(result);
}
transpose_to_labels(exec, result, &plan.target_labels)
}
fn eager_invalid_config(message: impl Into<String>) -> Error {
Error::invalid_argument(EAGER_EINSUM_OP, "configuration", message)
}
pub(crate) fn plan_subscripts(
subs: &Subscripts,
input_shapes: &[&[usize]],
) -> Result<ContractionTree> {
if input_shapes.is_empty() {
return Err(eager_invalid_config(
"eager einsum requires at least one input tensor",
));
}
if subs.inputs.len() != input_shapes.len() {
return Err(eager_invalid_config(format!(
"eager einsum subscripts expect {} inputs, got {}",
subs.inputs.len(),
input_shapes.len()
)));
}
ContractionTree::optimize(subs, input_shapes)
.map_err(|error| error.into_tensor_error(EAGER_EINSUM_OP))
}
fn take_labeled<'a>(
labeled: &mut [Option<LabeledTensor<'a>>],
index: usize,
role: &'static str,
) -> Result<LabeledTensor<'a>> {
labeled
.get_mut(index)
.ok_or_else(|| eager_invalid_config(format!("missing {role} operand at index {index}")))?
.take()
.ok_or_else(|| eager_invalid_config(format!("missing {role} operand at index {index}")))
}
fn find_label_axis(labels: &[u32], label: u32) -> Result<usize> {
labels
.iter()
.position(|candidate| *candidate == label)
.ok_or_else(|| eager_invalid_config(format!("label {label} missing from tensor labels")))
}
fn map_label_axes(source_labels: &[u32], target_labels: &[u32]) -> Result<Vec<usize>> {
map_label_occurrences(source_labels, target_labels).ok_or_else(|| {
eager_invalid_config(format!(
"cannot map label occurrences {source_labels:?} into {target_labels:?}"
))
})
}
fn label_size(label: u32, operands: &[&LabeledTensor<'_>]) -> Result<usize> {
for operand in operands {
if let Some(axis) = operand
.labels
.iter()
.position(|candidate| *candidate == label)
{
return Ok(operand.shape()[axis]);
}
}
Err(eager_invalid_config(format!(
"label {label} missing from eager einsum operands"
)))
}
fn reduce_tensor<'a>(
exec: &mut dyn BackendSession,
operand: LabeledTensor<'a>,
reduce_labels: &HashSet<u32>,
) -> Result<LabeledTensor<'a>> {
if reduce_labels.is_empty() {
return Ok(operand);
}
let reduce_axes: Vec<usize> = operand
.labels
.iter()
.enumerate()
.filter(|(_, label)| reduce_labels.contains(label))
.map(|(axis, _)| axis)
.collect();
if reduce_axes.is_empty() {
return Ok(operand);
}
let reduce_set: HashSet<usize> = reduce_axes.iter().copied().collect();
let labels: Vec<u32> = operand
.labels
.iter()
.enumerate()
.filter(|(axis, _)| !reduce_set.contains(axis))
.map(|(_, label)| *label)
.collect();
let operand_tensor = operand.tensor_owned(exec)?;
let tensor = exec.reduce_sum(&operand_tensor, &reduce_axes)?;
operand.reclaim_if_owned(exec);
Ok(LabeledTensor {
tensor: TensorValue::Owned(tensor),
labels,
})
}
fn diagonalize_repeated<'a>(
exec: &mut dyn BackendSession,
mut operand: LabeledTensor<'a>,
) -> Result<LabeledTensor<'a>> {
loop {
let mut seen = HashMap::new();
let mut repeated_pair = None;
for (axis, label) in operand.labels.iter().copied().enumerate() {
if let Some(first_axis) = seen.insert(label, axis) {
repeated_pair = Some((first_axis, axis));
break;
}
}
let Some((axis_a, axis_b)) = repeated_pair else {
return Ok(operand);
};
let operand_tensor = operand.tensor_owned(exec)?;
let tensor = exec.extract_diagonal(&operand_tensor, axis_a, axis_b)?;
let mut labels = operand.labels.clone();
labels.remove(axis_b);
operand.reclaim_if_owned(exec);
operand = LabeledTensor {
tensor: TensorValue::Owned(tensor),
labels,
};
}
}
fn embed_repeated<'a>(
exec: &mut dyn BackendSession,
mut operand: LabeledTensor<'a>,
output_labels: &[u32],
) -> Result<LabeledTensor<'a>> {
loop {
let mut embedded = false;
for &label in output_labels {
let current_count = operand
.labels
.iter()
.filter(|candidate| **candidate == label)
.count();
let output_count = output_labels
.iter()
.filter(|candidate| **candidate == label)
.count();
if output_count > current_count {
let axis_a = find_label_axis(&operand.labels, label)?;
let axis_b = axis_a + 1;
let operand_tensor = operand.tensor_owned(exec)?;
let tensor = exec.embed_diagonal(&operand_tensor, axis_a, axis_b)?;
let mut labels = operand.labels.clone();
labels.insert(axis_b, label);
operand.reclaim_if_owned(exec);
operand = LabeledTensor {
tensor: TensorValue::Owned(tensor),
labels,
};
embedded = true;
break;
}
}
if !embedded {
return Ok(operand);
}
}
}
fn transpose_to_labels<'a>(
exec: &mut dyn BackendSession,
operand: LabeledTensor<'a>,
target_labels: &[u32],
) -> Result<LabeledTensor<'a>> {
if operand.labels == target_labels {
return Ok(operand);
}
let perm = map_label_axes(target_labels, &operand.labels)?;
if perm
.iter()
.enumerate()
.all(|(axis, target)| axis == *target)
{
return Ok(operand);
}
let operand_tensor = operand.tensor_owned(exec)?;
let tensor = exec.transpose(&operand_tensor, &perm)?;
operand.reclaim_if_owned(exec);
Ok(LabeledTensor {
tensor: TensorValue::Owned(tensor),
labels: target_labels.to_vec(),
})
}
fn outer_product<'a>(
exec: &mut dyn BackendSession,
lhs: LabeledTensor<'a>,
rhs: LabeledTensor<'a>,
batch_labels: &[u32],
lhs_free_labels: &[u32],
rhs_free_labels: &[u32],
target_labels: Option<&[u32]>,
) -> Result<LabeledTensor<'a>> {
if lhs.labels == rhs.labels {
let tensor = exec.mul_read(lhs.tensor_read(), rhs.tensor_read())?;
let labels = lhs.labels.clone();
lhs.reclaim_if_owned(exec);
rhs.reclaim_if_owned(exec);
return Ok(LabeledTensor {
tensor: TensorValue::Owned(tensor),
labels,
});
}
let canonical_labels: Vec<u32> = lhs_free_labels
.iter()
.chain(rhs_free_labels.iter())
.chain(batch_labels.iter())
.copied()
.collect();
let combined_labels = select_outer_product_label_order(&canonical_labels, target_labels);
let combined_shape: Vec<usize> = combined_labels
.iter()
.map(|label| label_size(*label, &[&lhs, &rhs]))
.collect::<Result<_>>()?;
let lhs_dims = map_label_axes(&lhs.labels, &combined_labels)?;
let rhs_dims = map_label_axes(&rhs.labels, &combined_labels)?;
let tensor = match exec.execute_broadcast_multiply(
lhs.tensor.tensor_read(),
&combined_shape,
&lhs_dims,
rhs.tensor.tensor_read(),
&combined_shape,
&rhs_dims,
)? {
Some(tensor) => tensor,
None => match (
try_broadcast_tensor_read(&lhs.tensor, &combined_shape, &lhs_dims),
try_broadcast_tensor_read(&rhs.tensor, &combined_shape, &rhs_dims),
) {
(Some(lhs_read), Some(rhs_read)) => exec.mul_read(lhs_read?, rhs_read?)?,
_ => {
let lhs_input = lhs.tensor_owned(exec)?;
let rhs_input = rhs.tensor_owned(exec)?;
let lhs_tensor = exec.broadcast_in_dim(&lhs_input, &combined_shape, &lhs_dims)?;
let rhs_tensor = exec.broadcast_in_dim(&rhs_input, &combined_shape, &rhs_dims)?;
let tensor = exec.mul(&lhs_tensor, &rhs_tensor)?;
exec.reclaim_buffer(lhs_tensor);
exec.reclaim_buffer(rhs_tensor);
tensor
}
},
};
lhs.reclaim_if_owned(exec);
rhs.reclaim_if_owned(exec);
Ok(LabeledTensor {
tensor: TensorValue::Owned(tensor),
labels: combined_labels,
})
}
fn binary_contract<'a>(
exec: &mut dyn BackendSession,
lhs: LabeledTensor<'a>,
rhs: LabeledTensor<'a>,
survive_labels: &[u32],
reorder_result: bool,
) -> Result<LabeledTensor<'a>> {
if let Some(plan) = try_build_binary_dot_plan(&lhs.labels, &rhs.labels, survive_labels) {
return profile_eager_einsum_section("binary.fast_path", || {
execute_binary_dot_fast_plan(exec, lhs, rhs, plan, reorder_result)
});
}
let (survive_set, lhs_reduce, rhs_reduce) =
profile_eager_einsum_section("binary.pre_reduce_label_analysis", || {
let survive_set: HashSet<u32> = survive_labels.iter().copied().collect();
let rhs_label_set: HashSet<u32> = rhs.labels.iter().copied().collect();
let lhs_label_set: HashSet<u32> = lhs.labels.iter().copied().collect();
let lhs_reduce: HashSet<u32> = lhs
.labels
.iter()
.filter(|label| !rhs_label_set.contains(label) && !survive_set.contains(label))
.copied()
.collect();
let rhs_reduce: HashSet<u32> = rhs
.labels
.iter()
.filter(|label| !lhs_label_set.contains(label) && !survive_set.contains(label))
.copied()
.collect();
(survive_set, lhs_reduce, rhs_reduce)
});
let lhs = profile_eager_einsum_section("binary.reduce_lhs", || {
reduce_tensor(exec, lhs, &lhs_reduce)
})?;
let rhs = profile_eager_einsum_section("binary.reduce_rhs", || {
reduce_tensor(exec, rhs, &rhs_reduce)
})?;
let (batch_labels, contracting_labels, lhs_free_labels, rhs_free_labels) =
profile_eager_einsum_section("binary.classify_labels", || {
let lhs_label_set: HashSet<u32> = lhs.labels.iter().copied().collect();
let rhs_label_set: HashSet<u32> = rhs.labels.iter().copied().collect();
let mut batch_labels = Vec::new();
let mut contracting_labels = Vec::new();
let mut lhs_free_labels = Vec::new();
let mut rhs_free_labels = Vec::new();
for &label in &lhs.labels {
if rhs_label_set.contains(&label) {
if survive_set.contains(&label) {
if !batch_labels.contains(&label) {
batch_labels.push(label);
}
} else if !contracting_labels.contains(&label) {
contracting_labels.push(label);
}
} else if !lhs_free_labels.contains(&label) {
lhs_free_labels.push(label);
}
}
for &label in &rhs.labels {
if !lhs_label_set.contains(&label) && !rhs_free_labels.contains(&label) {
rhs_free_labels.push(label);
}
}
(
batch_labels,
contracting_labels,
lhs_free_labels,
rhs_free_labels,
)
});
let result = if contracting_labels.is_empty() {
outer_product(
exec,
lhs,
rhs,
&batch_labels,
&lhs_free_labels,
&rhs_free_labels,
reorder_result.then_some(survive_labels),
)?
} else {
let (labels, config) = profile_eager_einsum_section("binary.build_dot_config", || {
let lhs_contracting_dims: Vec<usize> = contracting_labels
.iter()
.map(|label| find_label_axis(&lhs.labels, *label))
.collect::<Result<_>>()?;
let rhs_contracting_dims: Vec<usize> = contracting_labels
.iter()
.map(|label| find_label_axis(&rhs.labels, *label))
.collect::<Result<_>>()?;
let lhs_batch_dims: Vec<usize> = batch_labels
.iter()
.map(|label| find_label_axis(&lhs.labels, *label))
.collect::<Result<_>>()?;
let rhs_batch_dims: Vec<usize> = batch_labels
.iter()
.map(|label| find_label_axis(&rhs.labels, *label))
.collect::<Result<_>>()?;
let labels: Vec<u32> = lhs_free_labels
.iter()
.chain(rhs_free_labels.iter())
.chain(batch_labels.iter())
.copied()
.collect();
let config = DotGeneralConfig {
lhs_contracting_dims,
rhs_contracting_dims,
lhs_batch_dims,
rhs_batch_dims,
};
Ok::<_, Error>((labels, config))
})?;
let tensor = profile_eager_einsum_section("binary.dot_general", || {
exec.dot_general_read(lhs.tensor_read(), rhs.tensor_read(), &config)
})?;
lhs.reclaim_if_owned(exec);
rhs.reclaim_if_owned(exec);
LabeledTensor {
tensor: TensorValue::Owned(tensor),
labels,
}
};
if !reorder_result {
return Ok(result);
}
let result_label_set: HashSet<u32> = result.labels.iter().copied().collect();
let target_labels: Vec<u32> = survive_labels
.iter()
.filter(|label| result_label_set.contains(label))
.copied()
.collect();
profile_eager_einsum_section("binary.final_transpose", || {
transpose_to_labels(exec, result, &target_labels)
})
}
fn eager_einsum_exec_values<'a>(
exec: &mut dyn BackendSession,
inputs: Vec<TensorValue<'a>>,
tree: &ContractionTree,
) -> Result<Tensor> {
record_eager_einsum_profile("exec_values.enter", Duration::ZERO);
let subscripts = &tree.subscripts;
let input_count = subscripts.inputs.len();
let output_labels = &subscripts.output;
let mut labeled: Vec<Option<LabeledTensor<'a>>> =
profile_eager_einsum_section("exec.init_labeled", || {
inputs
.into_iter()
.zip(subscripts.inputs.iter())
.map(|(tensor, labels)| {
Some(LabeledTensor {
tensor,
labels: labels.clone(),
})
})
.collect()
});
profile_eager_einsum_section("exec.diagonalize_inputs", || -> Result<()> {
for index in 0..labeled.len() {
let operand = take_labeled(&mut labeled, index, "input")?;
labeled[index] = Some(diagonalize_repeated(exec, operand)?);
}
Ok(())
})?;
if input_count == 1 || tree.step_count() == 0 {
let operand = take_labeled(&mut labeled, 0, "input")?;
let output_set: HashSet<u32> = output_labels.iter().copied().collect();
let reduce_labels: HashSet<u32> = operand
.labels
.iter()
.filter(|label| !output_set.contains(label))
.copied()
.collect();
let reduced = reduce_tensor(exec, operand, &reduce_labels)?;
let embedded = embed_repeated(exec, reduced, output_labels)?;
let reordered = transpose_to_labels(exec, embedded, output_labels)?;
return reordered.tensor.into_tensor(exec);
}
for step_idx in 0..tree.step_count() {
let (left, right) = tree.step_pair(step_idx).ok_or_else(|| {
eager_invalid_config(format!("missing contraction pair for step {step_idx}"))
})?;
let (_, _, step_output_labels) = tree.step_subscripts(step_idx).ok_or_else(|| {
eager_invalid_config(format!(
"missing contraction subscripts for step {step_idx}"
))
})?;
let lhs = take_labeled(&mut labeled, left, "lhs")?;
let rhs = take_labeled(&mut labeled, right, "rhs")?;
let result = profile_eager_einsum_section("exec.binary_contract", || {
binary_contract(
exec,
lhs,
rhs,
step_output_labels,
step_idx + 1 == tree.step_count(),
)
})?;
labeled.push(Some(result));
}
let final_index = input_count + tree.step_count() - 1;
let result = take_labeled(&mut labeled, final_index, "result")?;
let extra_labels: HashSet<u32> =
profile_eager_einsum_section("exec.final_label_analysis", || {
let output_set: HashSet<u32> = output_labels.iter().copied().collect();
result
.labels
.iter()
.filter(|label| !output_set.contains(label))
.copied()
.collect()
});
let reduced = profile_eager_einsum_section("exec.final_reduce", || {
reduce_tensor(exec, result, &extra_labels)
})?;
let reordered = profile_eager_einsum_section("exec.final_transpose", || {
transpose_to_labels(exec, reduced, output_labels)
})?;
reordered.tensor.into_tensor(exec)
}
#[cfg(all(feature = "autodiff", test))]
pub(crate) fn eager_einsum_with_tree(
exec: &mut dyn BackendSession,
inputs: &[&Tensor],
tree: &ContractionTree,
) -> Result<Tensor> {
eager_einsum_exec(exec, inputs, tree)
}
pub(crate) fn eager_einsum_exec(
exec: &mut dyn BackendSession,
inputs: &[&Tensor],
tree: &ContractionTree,
) -> Result<Tensor> {
record_eager_einsum_profile("exec.enter", Duration::ZERO);
let values = inputs
.iter()
.map(|tensor| TensorValue::Borrowed(tensor))
.collect();
eager_einsum_exec_values(exec, values, tree)
}
pub(crate) fn eager_einsum_exec_read(
exec: &mut dyn BackendSession,
inputs: &[TensorRead<'_>],
tree: &ContractionTree,
) -> Result<Tensor> {
record_eager_einsum_profile("exec_read.enter", Duration::ZERO);
let values = inputs
.iter()
.map(|input| match input {
TensorRead::Tensor(tensor) => TensorValue::Borrowed(tensor),
TensorRead::View(view) => TensorValue::View(view.clone()),
})
.collect();
eager_einsum_exec_values(exec, values, tree)
}
pub(crate) fn eager_einsum_exec_read_into(
exec: &mut dyn BackendSession,
inputs: &[TensorRead<'_>],
tree: &ContractionTree,
out: TensorWrite<'_>,
) -> Result<()> {
record_eager_einsum_profile("exec_read_into.enter", Duration::ZERO);
let subscripts = &tree.subscripts;
if inputs.len() == 2
&& subscripts.inputs.len() == 2
&& inputs[0].shape().len() == subscripts.inputs[0].len()
&& inputs[1].shape().len() == subscripts.inputs[1].len()
{
if let Some(plan) = try_build_binary_dot_plan(
&subscripts.inputs[0],
&subscripts.inputs[1],
&subscripts.output,
) {
if plan.result_labels == plan.target_labels {
return profile_eager_einsum_section("binary.fast_dot_general_into", || {
exec.dot_general_read_into(
inputs[0].clone(),
inputs[1].clone(),
&plan.config,
out,
)
});
}
}
}
let result = eager_einsum_exec_read(exec, inputs, tree)?;
exec.copy_read_into(TensorRead::from_tensor(&result), out)
}
pub(crate) fn eager_einsum_exec_read_into_accum(
exec: &mut dyn BackendSession,
inputs: &[TensorRead<'_>],
tree: &ContractionTree,
accumulation: DotGeneralAccumulation,
mut out: TensorWrite<'_>,
) -> Result<()> {
record_eager_einsum_profile("exec_read_into_accum.enter", Duration::ZERO);
let subscripts = &tree.subscripts;
if inputs.len() == 2
&& subscripts.inputs.len() == 2
&& inputs[0].shape().len() == subscripts.inputs[0].len()
&& inputs[1].shape().len() == subscripts.inputs[1].len()
{
if let Some(plan) = try_build_binary_dot_plan(
&subscripts.inputs[0],
&subscripts.inputs[1],
&subscripts.output,
) {
if plan.result_labels == plan.target_labels {
return profile_eager_einsum_section("binary.fast_dot_general_into_accum", || {
exec.dot_general_read_into_accum(
inputs[0].clone(),
inputs[1].clone(),
&plan.config,
accumulation,
out,
)
});
}
}
}
let result = eager_einsum_exec_read(exec, inputs, tree)?;
let accumulation = DotGeneralAccumulation {
lhs_conj: false,
rhs_conj: false,
..accumulation
};
tenferro_tensor::backend::accumulate_dot_result_into(&result, accumulation, &mut out)
}
fn tensor_value_from_read(input: TensorRead<'_>) -> TensorValue<'_> {
match input {
TensorRead::Tensor(tensor) => TensorValue::Borrowed(tensor),
TensorRead::View(view) => TensorValue::View(view),
}
}
fn eager_einsum_exec_binary_read_fast(
exec: &mut dyn BackendSession,
inputs: &[TensorRead<'_>],
subscripts: &Subscripts,
plan: BinaryDotPlan,
) -> Result<Tensor> {
let lhs = LabeledTensor {
tensor: tensor_value_from_read(inputs[0].clone()),
labels: subscripts.inputs[0].clone(),
};
let rhs = LabeledTensor {
tensor: tensor_value_from_read(inputs[1].clone()),
labels: subscripts.inputs[1].clone(),
};
execute_binary_dot_fast_plan(exec, lhs, rhs, plan, true)
.and_then(|result| result.tensor.into_tensor(exec))
}
fn try_eager_einsum_binary_read_fast(
ctx: &mut impl TensorBackend,
inputs: &[TensorRead<'_>],
subscripts: &Subscripts,
) -> Option<Result<Tensor>> {
let total_started = eager_einsum_profile_enabled().then(Instant::now);
if inputs.len() != 2 || subscripts.inputs.len() != 2 {
return None;
}
if inputs[0].shape().len() != subscripts.inputs[0].len()
|| inputs[1].shape().len() != subscripts.inputs[1].len()
{
return None;
}
let plan = profile_eager_einsum_section("fast.build_plan", || {
try_build_binary_dot_plan(
&subscripts.inputs[0],
&subscripts.inputs[1],
&subscripts.output,
)
})?;
let result = profile_eager_einsum_section("fast.with_backend_session", || {
ctx.with_backend_session(|exec| {
eager_einsum_exec_binary_read_fast(exec, inputs, subscripts, plan)
})
});
if let Some(started) = total_started {
record_eager_einsum_profile("total", started.elapsed());
maybe_print_eager_einsum_profile();
}
Some(result)
}
#[cfg(test)]
pub(crate) fn eager_einsum(
ctx: &mut impl TensorBackend,
inputs: &[&Tensor],
subscripts: &str,
) -> Result<Tensor> {
let subscripts =
Subscripts::parse(subscripts).map_err(|error| error.into_tensor_error(EAGER_EINSUM_OP))?;
eager_einsum_subscripts(ctx, inputs, &subscripts)
}
pub(crate) fn eager_einsum_subscripts(
ctx: &mut impl TensorBackend,
inputs: &[&Tensor],
subscripts: &Subscripts,
) -> Result<Tensor> {
if inputs.len() == 2 {
let read_inputs = [
TensorRead::from_tensor(inputs[0]),
TensorRead::from_tensor(inputs[1]),
];
if let Some(result) = try_eager_einsum_binary_read_fast(ctx, &read_inputs, subscripts) {
return result;
}
}
if eager_einsum_profile_enabled() {
let total_started = Instant::now();
let shapes = profile_eager_einsum_section("shape_collect", || {
inputs
.iter()
.map(|tensor| tensor.shape())
.collect::<Vec<_>>()
});
let tree = profile_eager_einsum_section("plan_subscripts", || {
plan_subscripts(subscripts, &shapes)
})?;
let result = profile_eager_einsum_section("with_backend_session", || {
ctx.with_backend_session(|exec| eager_einsum_exec(exec, inputs, &tree))
});
record_eager_einsum_profile("total", total_started.elapsed());
maybe_print_eager_einsum_profile();
return result;
}
if eager_einsum_trace_enabled() {
let total_started = Instant::now();
let started = Instant::now();
let shapes: Vec<&[usize]> = inputs.iter().map(|tensor| tensor.shape()).collect();
let shape_us = started.elapsed().as_secs_f64() * 1.0e6;
let started = Instant::now();
let tree = plan_subscripts(subscripts, &shapes)?;
let plan_us = started.elapsed().as_secs_f64() * 1.0e6;
let started = Instant::now();
let result = ctx.with_backend_session(|exec| eager_einsum_exec(exec, inputs, &tree));
let exec_us = started.elapsed().as_secs_f64() * 1.0e6;
let total_us = total_started.elapsed().as_secs_f64() * 1.0e6;
eprintln!(
"tenferro_eager_einsum_profile,input_count={},shapes={:?},steps={},shape_us={shape_us:.3},plan_us={plan_us:.3},exec_us={exec_us:.3},total_us={total_us:.3}",
inputs.len(),
shapes,
tree.step_count(),
);
return result;
}
let shapes: Vec<&[usize]> = inputs.iter().map(|tensor| tensor.shape()).collect();
let tree = plan_subscripts(subscripts, &shapes)?;
ctx.with_backend_session(|exec| eager_einsum_exec(exec, inputs, &tree))
}
#[cfg(feature = "autodiff")]
pub(crate) fn eager_einsum_subscripts_with_session(
exec: &mut dyn BackendSession,
inputs: &[&Tensor],
subscripts: &Subscripts,
) -> Result<Tensor> {
let shapes: Vec<&[usize]> = inputs.iter().map(|tensor| tensor.shape()).collect();
let tree = plan_subscripts(subscripts, &shapes)?;
eager_einsum_exec(exec, inputs, &tree)
}
pub(crate) fn eager_einsum_read_subscripts(
ctx: &mut impl TensorBackend,
inputs: &[TensorRead<'_>],
subscripts: &Subscripts,
) -> Result<Tensor> {
if let Some(result) = try_eager_einsum_binary_read_fast(ctx, inputs, subscripts) {
return result;
}
let shapes: Vec<&[usize]> = inputs.iter().map(TensorRead::shape).collect();
let tree = plan_subscripts(subscripts, &shapes)?;
ctx.with_backend_session(|exec| eager_einsum_exec_read(exec, inputs, &tree))
}
#[cfg(test)]
pub(crate) fn eager_einsum_owned(
ctx: &mut impl TensorBackend,
inputs: Vec<Tensor>,
subscripts: &str,
) -> Result<Tensor> {
let subscripts =
Subscripts::parse(subscripts).map_err(|error| error.into_tensor_error(EAGER_EINSUM_OP))?;
eager_einsum_owned_subscripts(ctx, inputs, &subscripts)
}
#[cfg(test)]
pub(crate) fn eager_einsum_owned_subscripts(
ctx: &mut impl TensorBackend,
inputs: Vec<Tensor>,
subscripts: &Subscripts,
) -> Result<Tensor> {
let shapes: Vec<&[usize]> = inputs.iter().map(|tensor| tensor.shape()).collect();
let tree = plan_subscripts(subscripts, &shapes)?;
let values = inputs.into_iter().map(TensorValue::Owned).collect();
ctx.with_backend_session(|exec| eager_einsum_exec_values(exec, values, &tree))
}
#[cfg(test)]
mod tests;