Skip to main content

tenferro_einsum/planning/
tree.rs

1use std::collections::{HashMap, HashSet};
2use std::fmt;
3use std::mem::{size_of, size_of_val};
4
5use omeco::{
6    CodeOptimizer, EinCode as OmecoEinCode, Initializer, NestedEinsum, ScoreFunction, TreeSA,
7};
8
9use crate::cache::{saturating_sum, vec_of_vec_retained_bytes, vec_retained_bytes};
10use crate::planning::plan::{compile_step_plans, DiagPlan, GemmPlan, ReducePlan, StepPlan};
11use crate::syntax::subscripts::Subscripts;
12use crate::util::{build_size_dict, intermediate_subs};
13use crate::{Error, Result};
14
15/// A single step in the contraction sequence.
16pub(crate) struct ContractionStep {
17    pub(crate) left: usize,
18    pub(crate) right: usize,
19}
20
21/// Public options for automatic contraction-path optimization.
22///
23/// The default planner uses TreeSA with a greedy initializer and zero annealing
24/// iterations. This keeps the public API on a single optimizer family while
25/// making the default behavior effectively "greedy-only".
26///
27/// # Determinism
28///
29/// While annealing is disabled (`niters == 0` or empty `betas`, as in the
30/// default), the planned path is a function of the subscripts and shapes
31/// alone: the same spec yields the same path in every process and on every
32/// call. With an annealing schedule the path comes from omeco's TreeSA, whose
33/// annealing is seeded but whose greedy initializer (omeco 0.2.6) breaks ties
34/// in `HashMap` order, so annealed paths are not guaranteed to agree across
35/// processes.
36#[derive(Debug, Clone)]
37pub struct ContractionOptimizerOptions {
38    /// Inverse-temperature schedule for TreeSA.
39    pub betas: Vec<f64>,
40    /// Number of independent TreeSA trials.
41    pub ntrials: usize,
42    /// Annealing iterations per temperature level.
43    pub niters: usize,
44    /// Score function used by TreeSA.
45    pub score: ScoreFunction,
46}
47
48impl Default for ContractionOptimizerOptions {
49    fn default() -> Self {
50        Self {
51            betas: Vec::new(),
52            ntrials: 1,
53            niters: 0,
54            score: ScoreFunction::default(),
55        }
56    }
57}
58
59impl ContractionOptimizerOptions {
60    fn to_treesa(&self) -> TreeSA {
61        TreeSA::new(
62            self.betas.clone(),
63            self.ntrials,
64            self.niters,
65            Initializer::Greedy,
66            self.score.clone(),
67        )
68    }
69
70    /// Whether TreeSA would run any annealing iteration.
71    fn anneals(&self) -> bool {
72        self.niters > 0 && !self.betas.is_empty()
73    }
74
75    pub(crate) fn validate(&self) -> Result<()> {
76        if self.ntrials == 0 {
77            return Err(Error::planning(
78                "contraction optimizer ntrials must be at least 1",
79            ));
80        }
81        if self.betas.iter().any(|value| value.is_nan()) {
82            return Err(Error::planning(
83                "contraction optimizer betas must not contain NaN",
84            ));
85        }
86        if self.score.tc_weight.is_nan()
87            || self.score.sc_weight.is_nan()
88            || self.score.rw_weight.is_nan()
89            || self.score.sc_target.is_nan()
90        {
91            return Err(Error::planning(
92                "contraction optimizer score fields must not contain NaN",
93            ));
94        }
95        Ok(())
96    }
97}
98
99/// Contraction tree determining pairwise contraction order for N-ary einsum.
100///
101/// When contracting more than two tensors, the order in which pairwise
102/// contractions are performed significantly affects performance.
103/// `ContractionTree` encodes this order as a binary tree.
104///
105/// # Optimization
106///
107/// Use [`ContractionTree::optimize`] for automatic cost-based optimization
108/// (e.g., greedy algorithm based on tensor sizes), or
109/// [`ContractionTree::from_pairs`] for manual specification.
110pub struct ContractionTree {
111    /// Original subscripts.
112    pub(crate) subscripts: Subscripts,
113    /// Steps in the contraction (empty for single-tensor case).
114    pub(crate) steps: Vec<ContractionStep>,
115    /// Label → dimension size mapping.
116    pub(crate) size_dict: HashMap<u32, usize>,
117    /// Subscripts for each operand (0..input_count from input, then intermediates).
118    pub(crate) operand_subs: Vec<Vec<u32>>,
119    /// Pre-compiled step plans (cached to avoid recomputation per execute call).
120    pub(crate) step_plans: Vec<StepPlan>,
121}
122
123impl fmt::Debug for ContractionTree {
124    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
125        f.debug_struct("ContractionTree")
126            .field("input_count", &self.subscripts.inputs.len())
127            .field("output_rank", &self.subscripts.output.len())
128            .field("steps_len", &self.steps.len())
129            .field("size_dict_len", &self.size_dict.len())
130            .field("operand_subs_len", &self.operand_subs.len())
131            .field("step_plans_len", &self.step_plans.len())
132            .finish_non_exhaustive()
133    }
134}
135
136impl ContractionTree {
137    /// Automatically compute an optimized contraction order.
138    ///
139    /// Uses a cost-based heuristic (greedy algorithm) to determine
140    /// the pairwise contraction sequence that minimizes total operation count.
141    /// The path is deterministic: for fixed subscripts and shapes it is the
142    /// same in every process (see [`ContractionOptimizerOptions`]).
143    ///
144    /// # Arguments
145    ///
146    /// * `subscripts` — Einsum subscripts for all tensors
147    /// * `shapes` — Shape of each input tensor
148    ///
149    /// # Examples
150    ///
151    /// ```rust
152    /// use tenferro_einsum::{ContractionTree, Subscripts};
153    ///
154    /// let subs = Subscripts::parse("abcdef,bf,cf,df,ef->f").unwrap();
155    /// let shapes = [&[2, 3, 2, 4, 3, 9][..], &[3, 9], &[2, 9], &[4, 9], &[3, 9]];
156    /// let tree = ContractionTree::optimize(&subs, &shapes).unwrap();
157    /// assert_eq!(tree.step_count(), 4);
158    /// // Equal-cost candidates are broken by operand index, never by hashing.
159    /// let again = ContractionTree::optimize(&subs, &shapes).unwrap();
160    /// for step in 0..tree.step_count() {
161    ///     assert_eq!(tree.step_pair(step), again.step_pair(step));
162    /// }
163    /// ```
164    ///
165    /// # Errors
166    ///
167    /// Returns [`Error::Validation`] when subscripts and shapes have different
168    /// ranks or incompatible dimensions, or [`Error::Planning`] when no valid
169    /// contraction order can be constructed.
170    pub fn optimize(subscripts: &Subscripts, shapes: &[&[usize]]) -> Result<Self> {
171        Self::optimize_with_options(subscripts, shapes, &ContractionOptimizerOptions::default())
172    }
173
174    /// Automatically compute an optimized contraction order with explicit
175    /// planner options.
176    ///
177    /// For three or more operands with an annealing schedule (`niters > 0`
178    /// and non-empty `betas`), this routes planning through TreeSA using the
179    /// provided configuration. Without annealing, including the default
180    /// options, TreeSA would return its greedy initializer unchanged, so the
181    /// deterministic greedy planner runs directly and the path is identical in
182    /// every process. One or two operands need no ordering search; their trees
183    /// are built directly after validating the options.
184    ///
185    /// # Errors
186    ///
187    /// Returns [`Error::Validation`] for rank, shape, or dimension mismatches,
188    /// or [`Error::Planning`] when planner options such as `ntrials` are
189    /// invalid or no contraction order can be constructed.
190    pub fn optimize_with_options(
191        subscripts: &Subscripts,
192        shapes: &[&[usize]],
193        options: &ContractionOptimizerOptions,
194    ) -> Result<Self> {
195        options.validate()?;
196        let input_count = subscripts.inputs.len();
197        if input_count <= 1 {
198            return Self::from_pairs(subscripts, shapes, &[]);
199        }
200        if input_count == 2 {
201            return Self::from_pairs(subscripts, shapes, &[(0, 1)]);
202        }
203
204        let size_dict = build_size_dict(subscripts, shapes, None)?;
205        // Without annealing TreeSA returns its greedy initializer unchanged,
206        // and omeco 0.2.6's greedy breaks ties in per-process `HashMap` order
207        // (#1963), so the deterministic greedy plans this case directly.
208        let pairs = if !options.anneals() {
209            optimize_self_greedy_pairs(subscripts, &size_dict)?
210        } else if let Some(omeco_pairs) = optimize_omeco_pairs(subscripts, &size_dict, options)? {
211            omeco_pairs
212        } else {
213            optimize_self_greedy_pairs(subscripts, &size_dict)?
214        };
215        Self::from_pairs(subscripts, shapes, &pairs)
216    }
217
218    /// Manually build a contraction tree from a pairwise contraction sequence.
219    ///
220    /// Each pair `(i, j)` specifies which two tensors (or intermediate results)
221    /// to contract next. Intermediate results are assigned indices starting
222    /// from the number of input tensors.
223    ///
224    /// # Arguments
225    ///
226    /// * `subscripts` — Einsum subscripts for all tensors
227    /// * `shapes` — Shape of each input tensor
228    /// * `pairs` — Ordered list of pairwise contractions
229    ///
230    /// # Examples
231    ///
232    /// ```rust
233    /// use tenferro_einsum::{ContractionTree, Subscripts};
234    ///
235    /// // Three tensors: A[ij] B[jk] C[kl] -> D[il]
236    /// // Contract B and C first, then A with the result:
237    /// let subs = Subscripts::new(&[&[0, 1], &[1, 2], &[2, 3]], &[0, 3]);
238    /// let shapes = [&[3, 4][..], &[4, 5], &[5, 6]];
239    /// let tree = ContractionTree::from_pairs(
240    ///     &subs,
241    ///     &shapes,
242    ///     &[(1, 2), (0, 3)],  // B*C -> T(index=3), then A*T -> D
243    /// ).unwrap();
244    /// ```
245    ///
246    /// # Errors
247    ///
248    /// Returns [`Error::Planning`] when the pair count, operand indices, or
249    /// intermediate sequence is invalid, or [`Error::Validation`] when the
250    /// supplied shapes do not match the subscripts.
251    pub fn from_pairs(
252        subscripts: &Subscripts,
253        shapes: &[&[usize]],
254        pairs: &[(usize, usize)],
255    ) -> Result<Self> {
256        let input_count = subscripts.inputs.len();
257        let required_steps = input_count.saturating_sub(1);
258        if pairs.len() != required_steps {
259            return Err(Error::planning(format!(
260                "explicit contraction path for {input_count} operands must have {required_steps} steps, got {}",
261                pairs.len()
262            )));
263        }
264        let size_dict = build_size_dict(subscripts, shapes, None)?;
265
266        let mut operand_subs: Vec<Vec<u32>> = subscripts.inputs.clone();
267        let mut live = vec![false; input_count + pairs.len()];
268        for slot in live.iter_mut().take(input_count) {
269            *slot = true;
270        }
271        let mut steps = Vec::new();
272
273        for (step_idx, &(left, right)) in pairs.iter().enumerate() {
274            let next_idx = input_count + step_idx;
275            if left == right {
276                return Err(Error::planning(format!(
277                    "pair ({left}, {right}) must reference two distinct live operands"
278                )));
279            }
280            if left >= next_idx || right >= next_idx {
281                return Err(Error::planning(format!(
282                    "pair ({left}, {right}) references non-existent operand"
283                )));
284            }
285            if !live[left] || !live[right] {
286                return Err(Error::planning(format!(
287                    "pair ({left}, {right}) references an operand or intermediate that is no longer live"
288                )));
289            }
290
291            // Labels needed by other live operands + final output
292            let mut needed: HashSet<u32> = subscripts.output.iter().copied().collect();
293            for (idx, subs) in operand_subs.iter().enumerate() {
294                if idx != left && idx != right && live[idx] {
295                    needed.extend(subs.iter().copied());
296                }
297            }
298
299            let new_subs = intermediate_subs(&operand_subs[left], &operand_subs[right], &needed);
300            operand_subs.push(new_subs);
301            live[left] = false;
302            live[right] = false;
303            live[next_idx] = true;
304            steps.push(ContractionStep { left, right });
305        }
306
307        let live_count = live.iter().filter(|&&is_live| is_live).count();
308        if live_count != 1 {
309            return Err(Error::planning(format!(
310                "explicit contraction path must leave exactly one live result, got {live_count}"
311            )));
312        }
313
314        let mut tree = Self {
315            subscripts: subscripts.clone(),
316            steps,
317            size_dict,
318            operand_subs,
319            step_plans: Vec::new(),
320        };
321        tree.step_plans = compile_step_plans(&tree)?;
322        Ok(tree)
323    }
324
325    /// Return the number of pairwise contraction steps in this tree.
326    ///
327    /// # Examples
328    ///
329    /// ```rust
330    /// use tenferro_einsum::{ContractionTree, Subscripts};
331    ///
332    /// let subs = Subscripts::new(&[&[0, 1], &[1, 2], &[2, 3]], &[0, 3]);
333    /// let tree = ContractionTree::from_pairs(
334    ///     &subs,
335    ///     &[&[2, 2], &[2, 2], &[2, 2]],
336    ///     &[(1, 2), (0, 3)],
337    /// )
338    /// .unwrap();
339    /// assert_eq!(tree.step_count(), 2);
340    /// ```
341    #[must_use]
342    pub fn step_count(&self) -> usize {
343        self.steps.len()
344    }
345
346    pub(crate) fn label_size(&self, label: u32) -> Option<usize> {
347        self.size_dict.get(&label).copied()
348    }
349
350    pub(crate) fn output_shape(&self) -> Vec<usize> {
351        self.subscripts
352            .output
353            .iter()
354            .filter_map(|label| self.label_size(*label))
355            .collect()
356    }
357
358    /// Return the operand indices for a pairwise contraction step.
359    ///
360    /// The returned indices refer to the original inputs (`0..input_count`) and
361    /// then to intermediates (`input_count..`) produced by earlier steps.
362    ///
363    /// # Examples
364    ///
365    /// ```rust
366    /// use tenferro_einsum::{ContractionTree, Subscripts};
367    ///
368    /// let subs = Subscripts::new(&[&[0, 1], &[1, 2], &[2, 3]], &[0, 3]);
369    /// let tree = ContractionTree::from_pairs(
370    ///     &subs,
371    ///     &[&[2, 2], &[2, 2], &[2, 2]],
372    ///     &[(1, 2), (0, 3)],
373    /// )
374    /// .unwrap();
375    /// assert_eq!(tree.step_pair(0), Some((1, 2)));
376    /// ```
377    #[must_use]
378    pub fn step_pair(&self, step_idx: usize) -> Option<(usize, usize)> {
379        self.steps.get(step_idx).map(|step| (step.left, step.right))
380    }
381
382    /// Return the `(lhs, rhs, output)` subscripts for a pairwise step.
383    ///
384    /// The output subscripts are the intermediate labels preserved after the
385    /// contraction, or the final output labels on the last step.
386    ///
387    /// # Examples
388    ///
389    /// ```rust
390    /// use tenferro_einsum::{ContractionTree, Subscripts};
391    ///
392    /// let subs = Subscripts::new(&[&[0, 1], &[1, 2], &[2, 3]], &[0, 3]);
393    /// let tree = ContractionTree::from_pairs(
394    ///     &subs,
395    ///     &[&[2, 2], &[2, 2], &[2, 2]],
396    ///     &[(1, 2), (0, 3)],
397    /// )
398    /// .unwrap();
399    /// let (lhs, rhs, out) = tree.step_subscripts(0).unwrap();
400    /// assert_eq!(lhs, &[1, 2]);
401    /// assert_eq!(rhs, &[2, 3]);
402    /// assert_eq!(out, &[1, 3]);
403    /// ```
404    #[must_use]
405    pub fn step_subscripts(&self, step_idx: usize) -> Option<(&[u32], &[u32], &[u32])> {
406        let input_count = self.subscripts.inputs.len();
407        let step = self.steps.get(step_idx)?;
408        let result_idx = input_count + step_idx;
409        let output_subs = if step_idx + 1 == self.steps.len() {
410            &self.subscripts.output
411        } else {
412            &self.operand_subs[result_idx]
413        };
414        Some((
415            &self.operand_subs[step.left],
416            &self.operand_subs[step.right],
417            output_subs,
418        ))
419    }
420
421    /// Return the precomputed lowering plan for one pairwise contraction step.
422    ///
423    /// # Examples
424    ///
425    /// ```rust
426    /// use tenferro_einsum::{ContractionTree, Subscripts};
427    ///
428    /// let subs = Subscripts::new(&[&[0, 1], &[1, 2]], &[0, 2]);
429    /// let tree = ContractionTree::from_pairs(&subs, &[&[2, 3], &[3, 4]], &[(0, 1)]).unwrap();
430    ///
431    /// assert_eq!(tree.step_plan(0).unwrap().gemm().m(), 2);
432    /// ```
433    #[must_use]
434    pub fn step_plan(&self, step_idx: usize) -> Option<crate::lowering::PairwiseStepPlan<'_>> {
435        self.step_plans
436            .get(step_idx)
437            .map(crate::lowering::PairwiseStepPlan::new)
438    }
439
440    #[doc(hidden)]
441    #[must_use]
442    pub(crate) fn retained_bytes_for_cache_stats(&self) -> usize {
443        saturating_sum([
444            size_of::<Self>(),
445            subscripts_retained_bytes(&self.subscripts),
446            self.steps
447                .capacity()
448                .saturating_mul(size_of::<ContractionStep>()),
449            self.size_dict
450                .capacity()
451                .saturating_mul(size_of::<u32>().saturating_add(size_of::<usize>())),
452            vec_of_vec_retained_bytes(&self.operand_subs),
453            self.step_plans
454                .capacity()
455                .saturating_mul(size_of::<StepPlan>()),
456            saturating_sum(self.step_plans.iter().map(step_plan_retained_bytes)),
457        ])
458    }
459}
460
461fn subscripts_retained_bytes(subscripts: &Subscripts) -> usize {
462    saturating_sum([
463        vec_of_vec_retained_bytes(&subscripts.inputs),
464        vec_retained_bytes(&subscripts.output),
465    ])
466}
467
468fn reduce_plan_retained_bytes(plan: &ReducePlan) -> usize {
469    saturating_sum([
470        vec_retained_bytes(&plan.original_subs),
471        vec_retained_bytes(&plan.kept_subs),
472        vec_retained_bytes(&plan.out_shape),
473    ])
474}
475
476fn diag_plan_retained_bytes(plan: &DiagPlan) -> usize {
477    saturating_sum([
478        vec_retained_bytes(&plan.stages),
479        saturating_sum(plan.stages.iter().map(|stage| {
480            saturating_sum([
481                vec_retained_bytes(&stage.axis_pairs),
482                vec_retained_bytes(&stage.result_subs),
483            ])
484        })),
485        vec_retained_bytes(&plan.result_subs),
486    ])
487}
488
489fn gemm_plan_retained_bytes(plan: &GemmPlan) -> usize {
490    saturating_sum([
491        plan.reduce_a.as_ref().map_or(0, reduce_plan_retained_bytes),
492        plan.reduce_b.as_ref().map_or(0, reduce_plan_retained_bytes),
493        vec_retained_bytes(&plan.subs_a),
494        vec_retained_bytes(&plan.subs_b),
495        vec_retained_bytes(&plan.lo_modes),
496        vec_retained_bytes(&plan.ro_modes),
497        vec_retained_bytes(&plan.sum_modes),
498        vec_retained_bytes(&plan.lo_sizes),
499        vec_retained_bytes(&plan.ro_sizes),
500        vec_retained_bytes(&plan.sum_sizes),
501        vec_retained_bytes(&plan.batch_sizes),
502        vec_retained_bytes(&plan.target_a),
503        vec_retained_bytes(&plan.target_b),
504        vec_retained_bytes(&plan.c_gemm_shape),
505        vec_retained_bytes(&plan.expanded_shape),
506        vec_retained_bytes(&plan.canonical_modes),
507        vec_retained_bytes(&plan.a_gemm_shape),
508        vec_retained_bytes(&plan.b_gemm_shape),
509    ])
510}
511
512fn step_plan_retained_bytes(plan: &StepPlan) -> usize {
513    saturating_sum([
514        plan.diag_a.as_ref().map_or(0, diag_plan_retained_bytes),
515        plan.diag_b.as_ref().map_or(0, diag_plan_retained_bytes),
516        plan.strict_binary.as_ref().map_or(0, size_of_val),
517        gemm_plan_retained_bytes(&plan.gemm),
518    ])
519}
520
521fn optimize_omeco_pairs(
522    subscripts: &Subscripts,
523    size_dict: &HashMap<u32, usize>,
524    options: &ContractionOptimizerOptions,
525) -> Result<Option<Vec<(usize, usize)>>> {
526    #[cfg(test)]
527    tests::OMECO_CALLS.with(|count| count.set(count.get() + 1));
528    let code = OmecoEinCode::new(subscripts.inputs.clone(), subscripts.output.clone());
529    let optimizer = options.to_treesa();
530    let Some(nested) = optimizer.optimize(&code, size_dict) else {
531        return Ok(None);
532    };
533
534    let mut next_operand = subscripts.inputs.len();
535    let mut pairs = Vec::with_capacity(subscripts.inputs.len().saturating_sub(1));
536    nested_to_pairs(&nested, &mut next_operand, &mut pairs)?;
537    Ok(Some(pairs))
538}
539
540fn nested_to_pairs(
541    nested: &NestedEinsum<u32>,
542    next_operand: &mut usize,
543    pairs: &mut Vec<(usize, usize)>,
544) -> Result<usize> {
545    match nested {
546        NestedEinsum::Leaf { tensor_index } => Ok(*tensor_index),
547        NestedEinsum::Node { args, .. } => {
548            if args.len() != 2 {
549                return Err(Error::planning(format!(
550                    "omeco returned non-binary contraction node with {} children",
551                    args.len()
552                )));
553            }
554            let left = nested_to_pairs(&args[0], next_operand, pairs)?;
555            let right = nested_to_pairs(&args[1], next_operand, pairs)?;
556            pairs.push((left, right));
557            let result_idx = *next_operand;
558            *next_operand += 1;
559            Ok(result_idx)
560        }
561    }
562}
563
564fn build_operand_label_sets(operand_subs: &[Vec<u32>]) -> Vec<HashSet<u32>> {
565    operand_subs
566        .iter()
567        .map(|subs| subs.iter().copied().collect())
568        .collect()
569}
570
571fn build_needed_label_counts(
572    output_subs: &[u32],
573    available: &[usize],
574    operand_label_sets: &[HashSet<u32>],
575) -> HashMap<u32, usize> {
576    let mut counts = HashMap::new();
577    for &label in output_subs {
578        counts.entry(label).or_insert(1);
579    }
580    for &idx in available {
581        add_labels_to_counts(&mut counts, &operand_label_sets[idx]);
582    }
583    counts
584}
585
586fn add_labels_to_counts(counts: &mut HashMap<u32, usize>, labels: &HashSet<u32>) {
587    for &label in labels {
588        *counts.entry(label).or_insert(0) += 1;
589    }
590}
591
592fn remove_labels_from_counts(counts: &mut HashMap<u32, usize>, labels: &HashSet<u32>) {
593    for &label in labels {
594        match counts.get(&label).copied() {
595            Some(1) => {
596                counts.remove(&label);
597            }
598            Some(count) => {
599                counts.insert(label, count - 1);
600            }
601            None => {}
602        }
603    }
604}
605
606fn candidate_label_is_needed(
607    label: u32,
608    left: usize,
609    right: usize,
610    operand_label_sets: &[HashSet<u32>],
611    needed_label_counts: &HashMap<u32, usize>,
612) -> bool {
613    let mut selected_count = 0;
614    if operand_label_sets[left].contains(&label) {
615        selected_count += 1;
616    }
617    if operand_label_sets[right].contains(&label) {
618        selected_count += 1;
619    }
620    needed_label_counts.get(&label).copied().unwrap_or(0) > selected_count
621}
622
623fn collect_candidate_intermediate_subs(
624    subs_left: &[u32],
625    subs_right: &[u32],
626    left: usize,
627    right: usize,
628    operand_label_sets: &[HashSet<u32>],
629    needed_label_counts: &HashMap<u32, usize>,
630    output: &mut Vec<u32>,
631) {
632    output.clear();
633    for &label in subs_left.iter().chain(subs_right.iter()) {
634        if candidate_label_is_needed(label, left, right, operand_label_sets, needed_label_counts)
635            && !output.contains(&label)
636        {
637            output.push(label);
638        }
639    }
640}
641
642#[derive(Clone, Copy)]
643struct CandidateCostContext<'a> {
644    operand_label_sets: &'a [HashSet<u32>],
645    needed_label_counts: &'a HashMap<u32, usize>,
646    size_dict: &'a HashMap<u32, usize>,
647}
648
649fn candidate_contraction_cost(
650    subs_left: &[u32],
651    subs_right: &[u32],
652    left: usize,
653    right: usize,
654    context: CandidateCostContext<'_>,
655    candidate_subs: &mut Vec<u32>,
656) -> Result<usize> {
657    collect_candidate_intermediate_subs(
658        subs_left,
659        subs_right,
660        left,
661        right,
662        context.operand_label_sets,
663        context.needed_label_counts,
664        candidate_subs,
665    );
666    let mut cost = 1usize;
667    for &label in candidate_subs.iter() {
668        let size = context.size_dict.get(&label).copied().ok_or_else(|| {
669            Error::planning(format!(
670                "unknown size for label {label} in contraction cost"
671            ))
672        })?;
673        cost = cost.saturating_mul(size);
674    }
675    Ok(cost.max(1))
676}
677
678/// Deterministic greedy contraction order.
679///
680/// Each step contracts the connected pair (sharing at least one label) whose
681/// result is smallest, the size counting every label still needed after the
682/// pair: by another operand or by the output. Ties go to the smallest
683/// `(left, right)` operand index pair, where an intermediate takes the next
684/// index after all existing operands. When no remaining operands share a
685/// label, the two smallest indices form an outer product.
686///
687/// This is the greedy rule of omeco's `tree_greedy` with `alpha = 0` and
688/// `temperature = 0` (omeco, MIT, <https://github.com/GiggleLiu/omeco>, which
689/// ports `OMEinsumContractionOrders.jl`), with the deterministic vertex and
690/// tie ordering adopted upstream after omeco 0.2.6. omeco 0.2.6 iterates a
691/// `HashMap` to break ties, so its path differed between processes (#1963);
692/// tenferro plans its default (non-annealing) path here instead, until an
693/// omeco release carries the fix (<https://github.com/GiggleLiu/omeco/issues/44>). Only the
694/// pairs whose operands changed are re-scored after each step, because the
695/// cost of any other pair is unaffected by the merge.
696fn optimize_self_greedy_pairs(
697    subscripts: &Subscripts,
698    size_dict: &HashMap<u32, usize>,
699) -> Result<Vec<(usize, usize)>> {
700    use std::cmp::Reverse;
701    use std::collections::{BTreeMap, BTreeSet, BinaryHeap};
702
703    #[cfg(test)]
704    tests::SELF_GREEDY_CALLS.with(|count| count.set(count.get() + 1));
705    let input_count = subscripts.inputs.len();
706    let available: Vec<usize> = (0..input_count).collect();
707    let mut operand_subs: Vec<Vec<u32>> = subscripts.inputs.clone();
708    let mut operand_label_sets = build_operand_label_sets(&operand_subs);
709    let mut needed_label_counts =
710        build_needed_label_counts(&subscripts.output, &available, &operand_label_sets);
711    let mut live: BTreeSet<usize> = available.into_iter().collect();
712    let mut label_owners: BTreeMap<u32, BTreeSet<usize>> = BTreeMap::new();
713    for (operand, labels) in operand_label_sets.iter().enumerate() {
714        for &label in labels {
715            label_owners.entry(label).or_default().insert(operand);
716        }
717    }
718
719    let mut candidate_subs = Vec::new();
720    let mut heap: BinaryHeap<Reverse<(usize, usize, usize)>> = BinaryHeap::new();
721    let mut score = |left: usize,
722                     right: usize,
723                     operand_subs: &[Vec<u32>],
724                     operand_label_sets: &[HashSet<u32>],
725                     needed_label_counts: &HashMap<u32, usize>|
726     -> Result<Reverse<(usize, usize, usize)>> {
727        let cost = candidate_contraction_cost(
728            &operand_subs[left],
729            &operand_subs[right],
730            left,
731            right,
732            CandidateCostContext {
733                operand_label_sets,
734                needed_label_counts,
735                size_dict,
736            },
737            &mut candidate_subs,
738        )?;
739        Ok(Reverse((cost, left, right)))
740    };
741
742    let mut initial_pairs = BTreeSet::new();
743    for owners in label_owners.values() {
744        for &left in owners {
745            for &right in owners.range(left + 1..) {
746                initial_pairs.insert((left, right));
747            }
748        }
749    }
750    for (left, right) in initial_pairs {
751        heap.push(score(
752            left,
753            right,
754            &operand_subs,
755            &operand_label_sets,
756            &needed_label_counts,
757        )?);
758    }
759
760    let mut pairs: Vec<(usize, usize)> = Vec::with_capacity(input_count.saturating_sub(1));
761    while live.len() > 1 {
762        let mut chosen = None;
763        while let Some(Reverse((_, left, right))) = heap.pop() {
764            if live.contains(&left) && live.contains(&right) {
765                chosen = Some((left, right));
766                break;
767            }
768        }
769        let (left, right) = match chosen {
770            Some(pair) => pair,
771            None => {
772                let mut smallest = live.iter().copied();
773                match (smallest.next(), smallest.next()) {
774                    (Some(left), Some(right)) => (left, right),
775                    _ => break,
776                }
777            }
778        };
779        pairs.push((left, right));
780
781        let mut new_subs = Vec::new();
782        collect_candidate_intermediate_subs(
783            &operand_subs[left],
784            &operand_subs[right],
785            left,
786            right,
787            &operand_label_sets,
788            &needed_label_counts,
789            &mut new_subs,
790        );
791        let new_idx = operand_subs.len();
792        let new_label_set: HashSet<u32> = new_subs.iter().copied().collect();
793        remove_labels_from_counts(&mut needed_label_counts, &operand_label_sets[left]);
794        remove_labels_from_counts(&mut needed_label_counts, &operand_label_sets[right]);
795        add_labels_to_counts(&mut needed_label_counts, &new_label_set);
796        for operand in [left, right] {
797            for label in &operand_label_sets[operand] {
798                if let Some(owners) = label_owners.get_mut(label) {
799                    owners.remove(&operand);
800                }
801            }
802        }
803        let mut neighbors = BTreeSet::new();
804        for &label in &new_subs {
805            let owners = label_owners.entry(label).or_default();
806            neighbors.extend(owners.iter().copied());
807            owners.insert(new_idx);
808        }
809        live.remove(&left);
810        live.remove(&right);
811        live.insert(new_idx);
812        operand_subs.push(new_subs);
813        operand_label_sets.push(new_label_set);
814        for neighbor in neighbors {
815            heap.push(score(
816                neighbor,
817                new_idx,
818                &operand_subs,
819                &operand_label_sets,
820                &needed_label_counts,
821            )?);
822        }
823    }
824
825    Ok(pairs)
826}
827
828#[cfg(test)]
829mod tests;