use std::collections::{HashMap, HashSet};
use crate::planning::classify::classify_modes;
pub(crate) use crate::planning::strict_binary::{
compile_strict_binary_lowering_step_plan, StrictBinaryLoweringPlan,
};
use crate::planning::tree::ContractionTree;
use crate::Result as EinsumResult;
#[derive(Debug, PartialEq, Eq)]
pub(crate) struct ReducePlan {
pub(crate) original_subs: Vec<u32>,
pub(crate) kept_subs: Vec<u32>,
pub(crate) out_shape: Vec<usize>,
}
#[derive(Debug, PartialEq, Eq)]
pub(crate) struct DiagStage {
pub(crate) axis_pairs: Vec<(usize, usize)>,
pub(crate) result_subs: Vec<u32>,
}
#[derive(Debug, PartialEq, Eq)]
pub(crate) struct DiagPlan {
pub(crate) stages: Vec<DiagStage>,
pub(crate) result_subs: Vec<u32>,
}
#[derive(Debug, PartialEq, Eq)]
pub(crate) struct GemmPlan {
pub(crate) reduce_a: Option<ReducePlan>,
pub(crate) reduce_b: Option<ReducePlan>,
pub(crate) subs_a: Vec<u32>,
pub(crate) subs_b: Vec<u32>,
pub(crate) lo_modes: Vec<u32>,
pub(crate) ro_modes: Vec<u32>,
pub(crate) sum_modes: Vec<u32>,
pub(crate) lo_sizes: Vec<usize>,
pub(crate) ro_sizes: Vec<usize>,
pub(crate) sum_sizes: Vec<usize>,
pub(crate) batch_sizes: Vec<usize>,
pub(crate) m: usize,
pub(crate) n: usize,
pub(crate) k: usize,
pub(crate) target_a: Vec<u32>,
pub(crate) target_b: Vec<u32>,
pub(crate) c_gemm_shape: Vec<usize>,
pub(crate) expanded_shape: Vec<usize>,
pub(crate) canonical_modes: Vec<u32>,
pub(crate) needs_final_permute: bool,
pub(crate) a_gemm_shape: Vec<usize>,
pub(crate) b_gemm_shape: Vec<usize>,
}
#[derive(Debug, PartialEq, Eq)]
pub(crate) struct StepPlan {
pub(crate) diag_a: Option<DiagPlan>,
pub(crate) diag_b: Option<DiagPlan>,
pub(crate) strict_binary: Option<StrictBinaryLoweringPlan>,
pub(crate) gemm: GemmPlan,
}
pub(crate) fn compute_reduce_plan(
subs_self: &[u32],
subs_other: &[u32],
subs_out: &[u32],
size_dict: &HashMap<u32, usize>,
) -> Option<ReducePlan> {
let other_set: HashSet<u32> = subs_other.iter().copied().collect();
let out_set: HashSet<u32> = subs_out.iter().copied().collect();
let mut has_reduction = false;
let mut kept_subs = Vec::with_capacity(subs_self.len());
for &label in subs_self {
if !other_set.contains(&label) && !out_set.contains(&label) {
has_reduction = true;
} else {
kept_subs.push(label);
}
}
if !has_reduction {
return None;
}
let out_shape: Vec<usize> = kept_subs.iter().map(|m| size_dict[m]).collect();
Some(ReducePlan {
original_subs: subs_self.to_vec(),
kept_subs,
out_shape,
})
}
pub(crate) fn compute_diag_plan_for_labels(
subs: &[u32],
labels_to_extract: &HashSet<u32>,
) -> Option<DiagPlan> {
fn label_positions(subs: &[u32]) -> HashMap<u32, Vec<usize>> {
let mut positions = HashMap::new();
for (i, &label) in subs.iter().enumerate() {
positions.entry(label).or_insert_with(Vec::new).push(i);
}
positions
}
fn build_diag_stage(subs: &[u32], labels_to_extract: &HashSet<u32>) -> Option<DiagStage> {
let positions_by_label = label_positions(subs);
let mut seen = HashSet::new();
let repeated_labels: Vec<u32> = subs
.iter()
.copied()
.filter(|label| {
labels_to_extract.contains(label)
&& positions_by_label
.get(label)
.is_some_and(|positions| positions.len() > 1)
&& seen.insert(*label)
})
.collect();
if repeated_labels.is_empty() {
return None;
}
let mut axis_pairs = Vec::new();
let mut used = vec![false; subs.len()];
let mut diag_labels = Vec::new();
for label in repeated_labels {
let positions = &positions_by_label[&label];
for chunk in positions.chunks(2) {
if let [left, right] = chunk {
axis_pairs.push((*left, *right));
used[*left] = true;
used[*right] = true;
diag_labels.push(label);
}
}
}
if axis_pairs.is_empty() {
return None;
}
let mut result_subs = Vec::with_capacity(subs.len() - axis_pairs.len());
for (i, &label) in subs.iter().enumerate() {
if !used[i] {
result_subs.push(label);
}
}
result_subs.extend(diag_labels);
Some(DiagStage {
axis_pairs,
result_subs,
})
}
let mut current_subs = subs.to_vec();
let mut stages = Vec::new();
while let Some(stage) = build_diag_stage(¤t_subs, labels_to_extract) {
current_subs = stage.result_subs.clone();
stages.push(stage);
}
if stages.is_empty() {
None
} else {
Some(DiagPlan {
stages,
result_subs: current_subs,
})
}
}
pub(crate) fn compute_diag_plan(subs: &[u32]) -> Option<DiagPlan> {
let mut repeated_labels = HashSet::new();
let mut label_positions: HashMap<u32, Vec<usize>> = HashMap::new();
for (i, &label) in subs.iter().enumerate() {
label_positions.entry(label).or_default().push(i);
}
for &label in subs {
if label_positions
.get(&label)
.is_some_and(|positions| positions.len() > 1)
{
repeated_labels.insert(label);
}
}
compute_diag_plan_for_labels(subs, &repeated_labels)
}
pub(crate) fn compile_pairwise_step_plan(
subs_a: &[u32],
subs_b: &[u32],
subs_c: &[u32],
size_dict: &HashMap<u32, usize>,
) -> EinsumResult<StepPlan> {
let diag_a = compute_diag_plan(subs_a);
let diag_b = compute_diag_plan(subs_b);
let eff_subs_a = diag_a
.as_ref()
.map(|d| d.result_subs.as_slice())
.unwrap_or(subs_a);
let eff_subs_b = diag_b
.as_ref()
.map(|d| d.result_subs.as_slice())
.unwrap_or(subs_b);
let reduce_a = compute_reduce_plan(eff_subs_a, eff_subs_b, subs_c, size_dict);
let reduce_b = compute_reduce_plan(eff_subs_b, eff_subs_a, subs_c, size_dict);
let effective_a = reduce_a
.as_ref()
.map(|r| r.kept_subs.clone())
.unwrap_or_else(|| eff_subs_a.to_vec());
let effective_b = reduce_b
.as_ref()
.map(|r| r.kept_subs.clone())
.unwrap_or_else(|| eff_subs_b.to_vec());
let (batch_modes, lo_modes, ro_modes, sum_modes) =
classify_modes(&effective_a, &effective_b, subs_c);
let batch_sizes: Vec<usize> = batch_modes.iter().map(|m| size_dict[m]).collect();
let lo_sizes: Vec<usize> = lo_modes.iter().map(|m| size_dict[m]).collect();
let ro_sizes: Vec<usize> = ro_modes.iter().map(|m| size_dict[m]).collect();
let sum_sizes: Vec<usize> = sum_modes.iter().map(|m| size_dict[m]).collect();
let m = product_or_one_for_empty(&lo_sizes, "left-only")?;
let n = product_or_one_for_empty(&ro_sizes, "right-only")?;
let k = product_or_one_for_empty(&sum_sizes, "contracted")?;
let target_a: Vec<u32> = lo_modes
.iter()
.chain(sum_modes.iter())
.chain(batch_modes.iter())
.copied()
.collect();
let target_b: Vec<u32> = sum_modes
.iter()
.chain(ro_modes.iter())
.chain(batch_modes.iter())
.copied()
.collect();
let a_gemm_shape: Vec<usize> = std::iter::once(m)
.chain(std::iter::once(k))
.chain(batch_sizes.iter().copied())
.collect();
let b_gemm_shape: Vec<usize> = std::iter::once(k)
.chain(std::iter::once(n))
.chain(batch_sizes.iter().copied())
.collect();
let c_gemm_shape: Vec<usize> = std::iter::once(m)
.chain(std::iter::once(n))
.chain(batch_sizes.iter().copied())
.collect();
let expanded_shape: Vec<usize> = lo_sizes
.iter()
.chain(ro_sizes.iter())
.chain(batch_sizes.iter())
.copied()
.collect();
let canonical_modes: Vec<u32> = lo_modes
.iter()
.chain(ro_modes.iter())
.chain(batch_modes.iter())
.copied()
.collect();
let needs_final_permute = canonical_modes.as_slice() != subs_c;
let strict_binary =
compile_strict_binary_lowering_step_plan(subs_a, subs_b, subs_c, size_dict)?;
Ok(StepPlan {
diag_a,
diag_b,
strict_binary,
gemm: GemmPlan {
reduce_a,
reduce_b,
subs_a: effective_a,
subs_b: effective_b,
lo_modes,
ro_modes,
sum_modes,
lo_sizes,
ro_sizes,
sum_sizes,
batch_sizes,
m,
n,
k,
target_a,
target_b,
c_gemm_shape,
expanded_shape,
canonical_modes,
needs_final_permute,
a_gemm_shape,
b_gemm_shape,
},
})
}
fn product_or_one_for_empty(sizes: &[usize], label: &'static str) -> EinsumResult<usize> {
if sizes.is_empty() {
Ok(1)
} else {
checked_product(sizes, label)
}
}
fn checked_product(sizes: &[usize], label: &'static str) -> EinsumResult<usize> {
sizes.iter().try_fold(1usize, |acc, &size| {
acc.checked_mul(size).ok_or_else(|| {
crate::Error::InvalidArgument(format!(
"dimension product overflow while fusing {label} dimensions {sizes:?}"
))
})
})
}
pub(crate) fn compile_step_plans(tree: &ContractionTree) -> EinsumResult<Vec<StepPlan>> {
let input_count = tree.subscripts.inputs.len();
let size_dict = &tree.size_dict;
tree.steps
.iter()
.enumerate()
.map(|(step_idx, step)| {
let subs_a = &tree.operand_subs[step.left];
let subs_b = &tree.operand_subs[step.right];
let is_last = step_idx == tree.steps.len() - 1;
let subs_c = if is_last {
&tree.subscripts.output
} else {
&tree.operand_subs[input_count + step_idx]
};
compile_pairwise_step_plan(subs_a, subs_b, subs_c, size_dict)
})
.collect()
}
#[cfg(test)]
#[path = "plan_tests.rs"]
mod tests;