ocas_atom/tensor/
dummy.rs1use std::collections::HashMap;
7
8use crate::{Atom, AtomArena, AtomNode, Symbol};
9
10use super::spec::TensorRegistry;
11
12#[derive(Debug, Clone, PartialEq, Eq)]
14pub enum DummyError {
15 OverContracted(Symbol),
17 BadContraction(Symbol),
19}
20
21pub fn refresh_dummies<'a>(
28 ctx: &'a AtomArena<'a>,
29 expr: Atom<'a>,
30 registry: &TensorRegistry,
31) -> Result<Atom<'a>, DummyError> {
32 let mut index_counts: HashMap<Atom<'a>, usize> = HashMap::new();
35 collect_counts(expr, &mut index_counts);
36
37 let dummies: Vec<Atom<'a>> = index_counts
39 .iter()
40 .filter(|(_, c)| **c == 2)
41 .map(|(l, _)| *l)
42 .collect();
43
44 if dummies.is_empty() {
45 return Ok(expr);
46 }
47
48 let mut group_counters: HashMap<u64, usize> = HashMap::new();
50 let mut renames: HashMap<Atom<'a>, Atom<'a>> = HashMap::new();
51
52 for d in &dummies {
53 let sym = Symbol::new(&d.to_string());
54 let group = registry.index_group(sym);
55 let cnt = group_counters.entry(group).or_insert(0);
56 let new_label = if group == 0 {
57 ctx.var(&format!("d{}", cnt))
58 } else {
59 ctx.var(&format!("d{}_{}", group, cnt))
60 };
61 *cnt += 1;
62 renames.insert(*d, new_label);
63 }
64
65 Ok(rename_in_expr(ctx, expr, &renames))
66}
67
68fn collect_counts<'a>(atom: Atom<'a>, counts: &mut HashMap<Atom<'a>, usize>) {
69 match atom.node() {
70 AtomNode::Num(_) | AtomNode::Var(_) => {}
71 AtomNode::Fun(_, args) => {
72 for a in args.iter() {
73 *counts.entry(*a).or_insert(0) += 1;
74 }
75 }
76 AtomNode::Add(args) | AtomNode::Mul(args) => {
77 for a in args.iter() {
78 collect_counts(*a, counts);
79 }
80 }
81 AtomNode::Pow(base, exp) => {
82 collect_counts(*base, counts);
83 collect_counts(*exp, counts);
84 }
85 }
86}
87
88fn rename_in_expr<'a>(
89 ctx: &'a AtomArena<'a>,
90 atom: Atom<'a>,
91 renames: &HashMap<Atom<'a>, Atom<'a>>,
92) -> Atom<'a> {
93 match atom.node() {
94 AtomNode::Num(_) => atom,
95 AtomNode::Var(_) => renames.get(&atom).copied().unwrap_or(atom),
96 AtomNode::Fun(name, args) => {
97 let new_args: Vec<Atom<'a>> = args
98 .iter()
99 .map(|a| rename_in_expr(ctx, *a, renames))
100 .collect();
101 ctx.fun(name.as_str(), &new_args)
102 }
103 AtomNode::Add(args) => {
104 let new_args: Vec<Atom<'a>> = args
105 .iter()
106 .map(|a| rename_in_expr(ctx, *a, renames))
107 .collect();
108 ctx.add(&new_args)
109 }
110 AtomNode::Mul(args) => {
111 let new_args: Vec<Atom<'a>> = args
112 .iter()
113 .map(|a| rename_in_expr(ctx, *a, renames))
114 .collect();
115 ctx.mul(&new_args)
116 }
117 AtomNode::Pow(base, exp) => {
118 let new_base = rename_in_expr(ctx, *base, renames);
119 let new_exp = rename_in_expr(ctx, *exp, renames);
120 ctx.pow(new_base, new_exp)
121 }
122 }
123}
124
125#[cfg(test)]
126mod tests {
127 use super::*;
128 use crate::AtomArena;
129 use crate::Symbol;
130 use crate::tensor::spec::SymmetrySpec;
131 use ocas_core::arena::Arena;
132
133 #[test]
134 fn refresh_renames_dummy() {
135 let arena = Arena::new();
136 let ctx = AtomArena::new(&arena);
137 let mut reg = TensorRegistry::new();
138 reg.register(Symbol::new("T"), SymmetrySpec::none());
139 reg.register(Symbol::new("U"), SymmetrySpec::none());
140
141 let i = ctx.var("i");
142 let j = ctx.var("j");
143 let t = ctx.fun("T", &[i, j]);
144 let u = ctx.fun("U", &[j, i]);
145 let prod = ctx.mul(&[t, u]);
146 let result = refresh_dummies(&ctx, prod, ®).unwrap();
147 let s = result.to_string();
148 assert!(s.contains("d0"), "expected dummy d0 in: {s}");
150 }
151}