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
15pub(crate) struct ContractionStep {
17 pub(crate) left: usize,
18 pub(crate) right: usize,
19}
20
21#[derive(Debug, Clone)]
37pub struct ContractionOptimizerOptions {
38 pub betas: Vec<f64>,
40 pub ntrials: usize,
42 pub niters: usize,
44 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 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
99pub struct ContractionTree {
111 pub(crate) subscripts: Subscripts,
113 pub(crate) steps: Vec<ContractionStep>,
115 pub(crate) size_dict: HashMap<u32, usize>,
117 pub(crate) operand_subs: Vec<Vec<u32>>,
119 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 pub fn optimize(subscripts: &Subscripts, shapes: &[&[usize]]) -> Result<Self> {
171 Self::optimize_with_options(subscripts, shapes, &ContractionOptimizerOptions::default())
172 }
173
174 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 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 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 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 #[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 #[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 #[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 #[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
678fn 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;