Skip to main content

ocas_atom/tensor/
dummy.rs

1//! Dummy index management for tensor canonicalisation.
2//!
3//! Provides dummy-index refresh (rename to avoid conflicts), n-ary
4//! contraction normalization, and validation for tensor expressions.
5
6use std::collections::HashMap;
7
8use crate::{Atom, AtomArena, AtomNode, Symbol};
9
10use super::spec::TensorRegistry;
11
12/// Errors from dummy-index operations.
13#[derive(Debug, Clone, PartialEq, Eq)]
14pub enum DummyError {
15    /// An index label appeared more than twice in a tensor product.
16    OverContracted(Symbol),
17    /// Two slots with the same label have the same variance (not an upper/lower pair).
18    BadContraction(Symbol),
19}
20
21/// Refresh (rename) dummy indices in a tensor expression to avoid
22/// conflicts with external (free) indices.
23///
24/// Dummy indices (labels appearing exactly twice in a product, once upper
25/// and once lower) are replaced with fresh names from a per-group pool.
26/// External indices are left unchanged.
27pub fn refresh_dummies<'a>(
28    ctx: &'a AtomArena<'a>,
29    expr: Atom<'a>,
30    registry: &TensorRegistry,
31) -> Result<Atom<'a>, DummyError> {
32    // Collect index usage: label → (count, vec of (head, pos) references).
33    // For simplicity, handle only single-term products here.
34    let mut index_counts: HashMap<Atom<'a>, usize> = HashMap::new();
35    collect_counts(expr, &mut index_counts);
36
37    // Identify dummies (count == 2).
38    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    // Assign fresh names per group.
49    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, &reg).unwrap();
147        let s = result.to_string();
148        // i appears as external, j as dummy — j should be renamed to d0.
149        assert!(s.contains("d0"), "expected dummy d0 in: {s}");
150    }
151}