use std::collections::{HashMap, HashSet};
use computegraph::graph::GraphBuilder;
use computegraph::types::{OperationRole, ValueRef};
use smallvec::SmallVec;
use tenferro_ops::dim_expr::DimExpr;
use tenferro_ops::std_tensor_op::StdTensorOp;
use tenferro_tensor::DotGeneralConfig;
use crate::planning::tree::ContractionTree;
use crate::util::map_label_occurrences;
use crate::{Error, Result};
pub(crate) type AxisVec = SmallVec<[usize; 4]>;
#[derive(Clone, Debug)]
struct LabeledVal {
val: ValueRef<StdTensorOp>,
labels: Vec<u32>,
shape: Vec<DimExpr>,
}
fn localize_value_ref(
builder: &mut GraphBuilder<StdTensorOp>,
val: ValueRef<StdTensorOp>,
shape: &[DimExpr],
) -> ValueRef<StdTensorOp> {
match val {
ValueRef::Local(_) => val,
ValueRef::External(_) => {
let outputs = builder.add_operation(
StdTensorOp::Reshape {
to_shape: shape.to_vec(),
},
vec![val],
OperationRole::Primary,
);
ValueRef::Local(outputs[0])
}
}
}
fn builder_invalid_argument(message: impl Into<String>) -> Error {
Error::InvalidArgument(format!("einsum builder: {}", message.into()))
}
fn find_label_axis(labels: &[u32], label: u32) -> Result<usize> {
labels
.iter()
.position(|candidate| *candidate == label)
.ok_or_else(|| builder_invalid_argument(format!("missing label {label} in {labels:?}")))
}
fn map_label_axes(source_labels: &[u32], target_labels: &[u32]) -> Result<AxisVec> {
map_label_occurrences(source_labels, target_labels)
.map(|axes| axes.into_iter().collect())
.ok_or_else(|| {
builder_invalid_argument(format!(
"cannot map label occurrences {source_labels:?} into {target_labels:?}"
))
})
}
fn local_shape(rank: usize) -> Vec<DimExpr> {
DimExpr::input_shape(0, rank)
}
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()
}
fn labeled_operand<'a>(
operands: &'a [LabeledVal],
index: usize,
role: &'static str,
) -> Result<&'a LabeledVal> {
operands
.get(index)
.ok_or_else(|| builder_invalid_argument(format!("missing {role} operand at index {index}")))
}
fn reduce_val(
builder: &mut GraphBuilder<StdTensorOp>,
lv: &LabeledVal,
reduce_labels: &HashSet<u32>,
) -> LabeledVal {
if reduce_labels.is_empty() {
return lv.clone();
}
let reduce_axes: AxisVec = lv
.labels
.iter()
.enumerate()
.filter(|(_, l)| reduce_labels.contains(l))
.map(|(i, _)| i)
.collect();
if reduce_axes.is_empty() {
return lv.clone();
}
let reduce_set: HashSet<usize> = reduce_axes.iter().copied().collect();
let new_labels: Vec<u32> = lv
.labels
.iter()
.enumerate()
.filter(|(i, _)| !reduce_set.contains(i))
.map(|(_, &l)| l)
.collect();
let new_shape = local_shape(new_labels.len());
let outputs = builder.add_operation(
StdTensorOp::ReduceSum {
axes: reduce_axes.into_vec(),
},
vec![lv.val.clone()],
OperationRole::Primary,
);
LabeledVal {
val: ValueRef::Local(outputs[0]),
labels: new_labels,
shape: new_shape,
}
}
fn embed_repeated(
builder: &mut GraphBuilder<StdTensorOp>,
lv: &LabeledVal,
output_labels: &[u32],
) -> Result<LabeledVal> {
let mut result = lv.clone();
for &label in output_labels {
let current_count = result.labels.iter().filter(|&&l| l == label).count();
let output_count = output_labels.iter().filter(|&&l| l == label).count();
if output_count > current_count {
let axis_a = find_label_axis(&result.labels, label)?;
let axis_b = axis_a + 1;
let outputs = builder.add_operation(
StdTensorOp::EmbedDiag { axis_a, axis_b },
vec![result.val.clone()],
OperationRole::Primary,
);
let mut new_labels = result.labels.clone();
new_labels.insert(axis_b, label);
let new_shape = local_shape(new_labels.len());
result = LabeledVal {
val: ValueRef::Local(outputs[0]),
labels: new_labels,
shape: new_shape,
};
return embed_repeated(builder, &result, output_labels);
}
}
Ok(result)
}
fn diagonalize_repeated(builder: &mut GraphBuilder<StdTensorOp>, lv: &LabeledVal) -> LabeledVal {
let mut seen: HashMap<u32, usize> = HashMap::new();
for (i, &label) in lv.labels.iter().enumerate() {
if let Some(&first) = seen.get(&label) {
let outputs = builder.add_operation(
StdTensorOp::ExtractDiag {
axis_a: first,
axis_b: i,
},
vec![lv.val.clone()],
OperationRole::Primary,
);
let mut new_labels = lv.labels.clone();
new_labels.remove(i);
let new_shape = local_shape(new_labels.len());
let result = LabeledVal {
val: ValueRef::Local(outputs[0]),
labels: new_labels,
shape: new_shape,
};
return diagonalize_repeated(builder, &result);
}
seen.insert(label, i);
}
lv.clone()
}
fn binary_contract(
builder: &mut GraphBuilder<StdTensorOp>,
lhs: &LabeledVal,
rhs: &LabeledVal,
survive_labels: &[u32],
reorder_result: bool,
) -> Result<LabeledVal> {
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(|l| !rhs_label_set.contains(l) && !survive_set.contains(l))
.copied()
.collect();
let rhs_reduce: HashSet<u32> = rhs
.labels
.iter()
.filter(|l| !lhs_label_set.contains(l) && !survive_set.contains(l))
.copied()
.collect();
let lhs = reduce_val(builder, lhs, &lhs_reduce);
let rhs = reduce_val(builder, rhs, &rhs_reduce);
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 &l in &lhs.labels {
if rhs_label_set.contains(&l) {
if survive_set.contains(&l) {
if !batch_labels.contains(&l) {
batch_labels.push(l);
}
} else if !contracting_labels.contains(&l) {
contracting_labels.push(l);
}
} else if !lhs_free_labels.contains(&l) {
lhs_free_labels.push(l);
}
}
for &l in &rhs.labels {
if !lhs_label_set.contains(&l) && !rhs_free_labels.contains(&l) {
rhs_free_labels.push(l);
}
}
let result = if !contracting_labels.is_empty() {
let lhs_contracting_dims: AxisVec = contracting_labels
.iter()
.map(|l| find_label_axis(&lhs.labels, *l))
.collect::<Result<_>>()?;
let rhs_contracting_dims: AxisVec = contracting_labels
.iter()
.map(|l| find_label_axis(&rhs.labels, *l))
.collect::<Result<_>>()?;
let lhs_batch_dims: AxisVec = batch_labels
.iter()
.map(|l| find_label_axis(&lhs.labels, *l))
.collect::<Result<_>>()?;
let rhs_batch_dims: AxisVec = batch_labels
.iter()
.map(|l| find_label_axis(&rhs.labels, *l))
.collect::<Result<_>>()?;
let config = DotGeneralConfig {
lhs_contracting_dims: lhs_contracting_dims.into_vec(),
rhs_contracting_dims: rhs_contracting_dims.into_vec(),
lhs_batch_dims: lhs_batch_dims.into_vec(),
rhs_batch_dims: rhs_batch_dims.into_vec(),
};
let result_labels: Vec<u32> = lhs_free_labels
.iter()
.chain(rhs_free_labels.iter())
.chain(batch_labels.iter())
.copied()
.collect();
let result_shape = local_shape(result_labels.len());
let outputs = builder.add_operation(
StdTensorOp::DotGeneral { config },
vec![lhs.val.clone(), rhs.val.clone()],
OperationRole::Primary,
);
LabeledVal {
val: ValueRef::Local(outputs[0]),
labels: result_labels,
shape: result_shape,
}
} else {
outer_product(
builder,
&lhs,
&rhs,
&batch_labels,
&lhs_free_labels,
&rhs_free_labels,
reorder_result.then_some(survive_labels),
)?
};
if !reorder_result {
return Ok(result);
}
let current_labels = &result.labels;
if current_labels.is_empty() {
return Ok(result);
}
let result_label_set: HashSet<u32> = current_labels.iter().copied().collect();
let target_labels: Vec<u32> = survive_labels
.iter()
.filter(|l| result_label_set.contains(l))
.copied()
.collect();
if current_labels.len() == target_labels.len() && *current_labels == target_labels {
return Ok(result);
}
let perm = map_label_axes(&target_labels, current_labels)?;
if perm.iter().enumerate().all(|(i, &p)| i == p) {
return Ok(result);
}
let new_shape = local_shape(target_labels.len());
let outputs = builder.add_operation(
StdTensorOp::Transpose {
perm: perm.into_vec(),
},
vec![result.val.clone()],
OperationRole::Primary,
);
Ok(LabeledVal {
val: ValueRef::Local(outputs[0]),
labels: target_labels,
shape: new_shape,
})
}
fn outer_product(
builder: &mut GraphBuilder<StdTensorOp>,
lhs: &LabeledVal,
rhs: &LabeledVal,
batch_labels: &[u32],
lhs_free_labels: &[u32],
rhs_free_labels: &[u32],
target_labels: Option<&[u32]>,
) -> Result<LabeledVal> {
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);
if lhs.labels == rhs.labels {
let outputs = builder.add_operation(
StdTensorOp::Mul,
vec![lhs.val.clone(), rhs.val.clone()],
OperationRole::Primary,
);
return Ok(LabeledVal {
val: ValueRef::Local(outputs[0]),
labels: lhs.labels.clone(),
shape: lhs.shape.clone(),
});
}
let lhs_dims = map_label_axes(&lhs.labels, &combined_labels)?;
let rhs_dims = map_label_axes(&rhs.labels, &combined_labels)?;
let lhs_shape =
combined_shape_for_broadcast(&combined_labels, lhs, rhs, BroadcastPrimary::Lhs)?;
let rhs_shape =
combined_shape_for_broadcast(&combined_labels, lhs, rhs, BroadcastPrimary::Rhs)?;
let lhs_bc = builder.add_operation(
StdTensorOp::BroadcastInDim {
shape: lhs_shape.clone(),
dims: lhs_dims.into_vec(),
},
broadcast_inputs(lhs.val.clone(), rhs.val.clone(), &lhs_shape),
OperationRole::Primary,
);
let rhs_bc = builder.add_operation(
StdTensorOp::BroadcastInDim {
shape: rhs_shape.clone(),
dims: rhs_dims.into_vec(),
},
broadcast_inputs(rhs.val.clone(), lhs.val.clone(), &rhs_shape),
OperationRole::Primary,
);
let outputs = builder.add_operation(
StdTensorOp::Mul,
vec![ValueRef::Local(lhs_bc[0]), ValueRef::Local(rhs_bc[0])],
OperationRole::Primary,
);
let combined_rank = combined_labels.len();
Ok(LabeledVal {
val: ValueRef::Local(outputs[0]),
labels: combined_labels,
shape: local_shape(combined_rank),
})
}
#[derive(Clone, Copy)]
enum BroadcastPrimary {
Lhs,
Rhs,
}
fn combined_shape_for_broadcast(
combined_labels: &[u32],
lhs: &LabeledVal,
rhs: &LabeledVal,
primary: BroadcastPrimary,
) -> Result<Vec<DimExpr>> {
combined_labels
.iter()
.map(|&label| {
if let Some(axis) = lhs.labels.iter().position(|candidate| *candidate == label) {
let input_idx = match primary {
BroadcastPrimary::Lhs => 0,
BroadcastPrimary::Rhs => 1,
};
return Ok(DimExpr::InputDim { input_idx, axis });
}
if let Some(axis) = rhs.labels.iter().position(|candidate| *candidate == label) {
let input_idx = match primary {
BroadcastPrimary::Lhs => 1,
BroadcastPrimary::Rhs => 0,
};
return Ok(DimExpr::InputDim { input_idx, axis });
}
Err(builder_invalid_argument(format!(
"missing label {label} while building broadcast shape"
)))
})
.collect()
}
fn broadcast_inputs(
primary: ValueRef<StdTensorOp>,
secondary: ValueRef<StdTensorOp>,
shape: &[DimExpr],
) -> Vec<ValueRef<StdTensorOp>> {
let mut inputs = vec![primary];
if shape_uses_input(shape, 1) {
inputs.push(secondary);
}
inputs
}
fn shape_uses_input(shape: &[DimExpr], input_idx: usize) -> bool {
shape.iter().any(|dim| dim_expr_uses_input(dim, input_idx))
}
fn dim_expr_uses_input(dim: &DimExpr, input_idx: usize) -> bool {
match dim {
DimExpr::Const(_) => false,
DimExpr::InputDim {
input_idx: actual, ..
} => *actual == input_idx,
DimExpr::Add(lhs, rhs)
| DimExpr::Sub(lhs, rhs)
| DimExpr::Mul(lhs, rhs)
| DimExpr::FloorDiv(lhs, rhs)
| DimExpr::Min(lhs, rhs)
| DimExpr::Max(lhs, rhs) => {
dim_expr_uses_input(lhs, input_idx) || dim_expr_uses_input(rhs, input_idx)
}
}
}
pub(crate) fn build_einsum_graph(
builder: &mut GraphBuilder<StdTensorOp>,
tree: &ContractionTree,
input_vals: &[ValueRef<StdTensorOp>],
input_shapes: &[Vec<usize>],
) -> Result<ValueRef<StdTensorOp>> {
let input_shapes: Vec<Vec<DimExpr>> = input_shapes
.iter()
.map(|shape| DimExpr::from_concrete(shape))
.collect();
build_einsum_graph_dim_expr(builder, tree, input_vals, &input_shapes)
}
pub(crate) fn build_einsum_graph_dim_expr(
builder: &mut GraphBuilder<StdTensorOp>,
tree: &ContractionTree,
input_vals: &[ValueRef<StdTensorOp>],
input_shapes: &[Vec<DimExpr>],
) -> Result<ValueRef<StdTensorOp>> {
let subscripts = &tree.subscripts;
let input_count = subscripts.inputs.len();
if input_count != input_vals.len() {
return Err(builder_invalid_argument(format!(
"number of subscripts inputs ({input_count}) must match number of input values ({})",
input_vals.len()
)));
}
if input_vals.len() != input_shapes.len() {
return Err(builder_invalid_argument(format!(
"number of input values ({}) must match number of input shapes ({})",
input_vals.len(),
input_shapes.len()
)));
}
let output_labels = &subscripts.output;
let mut labeled: Vec<LabeledVal> = input_vals
.iter()
.zip(subscripts.inputs.iter())
.zip(input_shapes.iter())
.map(|((val, labels), shape)| {
if labels.len() != shape.len() {
return Err(builder_invalid_argument(format!(
"labels length ({}) must match shape rank ({})",
labels.len(),
shape.len()
)));
}
Ok(LabeledVal {
val: val.clone(),
labels: labels.clone(),
shape: local_shape(shape.len()),
})
})
.collect::<Result<_>>()?;
for lv in &mut labeled {
*lv = diagonalize_repeated(builder, lv);
}
if input_count == 1 || tree.step_count() == 0 {
let lv = &labeled[0];
let output_set: HashSet<u32> = output_labels.iter().copied().collect();
let reduce_labels: HashSet<u32> = lv
.labels
.iter()
.filter(|l| !output_set.contains(l))
.copied()
.collect();
let result = reduce_val(builder, lv, &reduce_labels);
let result = embed_repeated(builder, &result, output_labels)?;
if result.labels == *output_labels {
return Ok(localize_value_ref(builder, result.val, &result.shape));
}
let perm = map_label_axes(output_labels, &result.labels)?;
if perm.iter().enumerate().all(|(i, &p)| i == p) {
return Ok(localize_value_ref(builder, result.val, &result.shape));
}
let outputs = builder.add_operation(
StdTensorOp::Transpose {
perm: perm.into_vec(),
},
vec![result.val],
OperationRole::Primary,
);
return Ok(ValueRef::Local(outputs[0]));
}
for step_idx in 0..tree.step_count() {
let (left, right) = tree.step_pair(step_idx).ok_or_else(|| {
builder_invalid_argument(format!("missing contraction pair for step {step_idx}"))
})?;
let (_, _, step_out_labels) = tree.step_subscripts(step_idx).ok_or_else(|| {
builder_invalid_argument(format!(
"missing contraction subscripts for step {step_idx}"
))
})?;
let is_last = step_idx + 1 == tree.step_count();
let result = binary_contract(
builder,
labeled_operand(&labeled, left, "left")?,
labeled_operand(&labeled, right, "right")?,
step_out_labels,
is_last,
)?;
labeled.push(result);
}
let final_idx = input_count + tree.step_count() - 1;
let result = labeled_operand(&labeled, final_idx, "final result")?;
let output_set: HashSet<u32> = output_labels.iter().copied().collect();
let extra_labels: HashSet<u32> = result
.labels
.iter()
.filter(|l| !output_set.contains(l))
.copied()
.collect();
let result = reduce_val(builder, result, &extra_labels);
if result.labels == *output_labels {
return Ok(localize_value_ref(builder, result.val, &result.shape));
}
if result.labels.is_empty() && output_labels.is_empty() {
return Ok(localize_value_ref(builder, result.val, &result.shape));
}
let perm = map_label_axes(output_labels, &result.labels)?;
if perm.iter().enumerate().all(|(i, &p)| i == p) {
return Ok(localize_value_ref(builder, result.val, &result.shape));
}
let outputs = builder.add_operation(
StdTensorOp::Transpose {
perm: perm.into_vec(),
},
vec![result.val.clone()],
OperationRole::Primary,
);
Ok(ValueRef::Local(outputs[0]))
}