Skip to main content

ocas_atom/tensor/
canon.rs

1//! Tensor expression canonicalisation via graph isomorphism.
2//!
3//! Encodes a tensor-product expression (Mul of Fun nodes) into a graph whose
4//! vertex colours represent tensor heads / index slots and whose edges
5//! represent argument positions and contractions.  The graph-isomorphism
6//! engine [`super::graph`] then computes a canonical labelling, and the
7//! result is reconstructed as a normalised tensor expression with renamed
8//! dummy indices and reordered symmetric slots.
9
10use std::collections::HashMap;
11
12use crate::{Atom, AtomArena, AtomNode, Symbol};
13
14use super::graph::{CanonicalForm, Graph};
15use super::spec::TensorRegistry;
16
17/// Error while canonicalising a tensor expression.
18#[derive(Debug, Clone, PartialEq, Eq)]
19pub enum TensorCanonError {
20    ContractedMoreThanOnce(Symbol),
21    BadContraction(Symbol),
22    NotATensor(Symbol),
23    InconsistentOpenIndices,
24    UnsupportedPower,
25}
26
27/// Result of canonicalising a tensor expression.
28#[derive(Debug, Clone)]
29pub struct CanonicalTensor<'a> {
30    pub canonical_form: Atom<'a>,
31    pub external_indices: Vec<Atom<'a>>,
32    pub dummy_indices: Vec<Atom<'a>>,
33}
34
35// =========================================================================
36// Public API
37// =========================================================================
38
39/// Canonicalise a tensor expression.
40pub fn canonicalize_tensors<'a>(
41    ctx: &'a AtomArena<'a>,
42    expr: Atom<'a>,
43    registry: &TensorRegistry,
44) -> Result<CanonicalTensor<'a>, TensorCanonError> {
45    match expr.node() {
46        AtomNode::Add(terms) => {
47            let mut canon_terms: Vec<Atom<'a>> = Vec::new();
48            let mut first_external: Option<Vec<Atom<'a>>> = None;
49            let mut all_dummies: Vec<Atom<'a>> = Vec::new();
50
51            for term in terms.iter() {
52                let ct = canonicalize_single_term(ctx, *term, registry)?;
53                match &first_external {
54                    None => first_external = Some(ct.external_indices.clone()),
55                    Some(ext) if *ext != ct.external_indices => {
56                        return Err(TensorCanonError::InconsistentOpenIndices);
57                    }
58                    _ => {}
59                }
60                all_dummies.extend(ct.dummy_indices);
61                canon_terms.push(ct.canonical_form);
62            }
63
64            let canonical_form = if canon_terms.len() == 1 {
65                canon_terms.pop().unwrap()
66            } else {
67                ctx.add(&canon_terms)
68            };
69            Ok(CanonicalTensor {
70                canonical_form,
71                external_indices: first_external.unwrap_or_default(),
72                dummy_indices: all_dummies,
73            })
74        }
75        _ => canonicalize_single_term(ctx, expr, registry),
76    }
77}
78
79fn canonicalize_single_term<'a>(
80    ctx: &'a AtomArena<'a>,
81    expr: Atom<'a>,
82    registry: &TensorRegistry,
83) -> Result<CanonicalTensor<'a>, TensorCanonError> {
84    // Fast path: single tensor with all-symmetric slots. Sort args
85    // alphabetically so the result is input-order-independent.
86    #[allow(clippy::collapsible_if)]
87    if let AtomNode::Fun(name, args) = expr.node() {
88        if let Some(spec) = registry.spec(*name) {
89            let all_symmetric = !spec.symmetric_subsets.is_empty()
90                && spec.antisymmetric_subsets.is_empty()
91                && (0..args.len()).all(|pos| spec.is_slot_hidden(pos));
92            if all_symmetric {
93                let mut sorted: Vec<Atom<'a>> = args.to_vec();
94                sorted.sort_by_key(|a| match a.node() {
95                    AtomNode::Var(s) => s.as_str().to_string(),
96                    _ => a.to_string(),
97                });
98                let result = ctx.fun(name.as_str(), &sorted);
99                return Ok(CanonicalTensor {
100                    canonical_form: result,
101                    external_indices: sorted,
102                    dummy_indices: Vec::new(),
103                });
104            }
105        }
106    }
107
108    let (g, head_nodes, slot_labels) = tensor_to_graph(ctx, expr, registry)?;
109    let cf = g.canonize();
110    reconstruct(ctx, &cf, &head_nodes, &slot_labels, registry)
111}
112
113// =========================================================================
114// Graph encoding
115// =========================================================================
116
117#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
118enum TgNode {
119    Head(u64),
120    Slot(u64),
121    Scalar(u64),
122}
123
124#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
125enum TgEdge {
126    HeadToSlot(usize, u8),
127    Contraction(u64),
128}
129
130#[derive(Debug, Clone)]
131#[allow(dead_code)]
132struct HeadInfo {
133    symbol: Symbol,
134    slot_count: usize,
135    head_v: usize,
136    slot_verts: Vec<usize>,
137}
138
139/// slot_vertex_index → original label Atom for all slots.
140#[allow(clippy::type_complexity)]
141fn tensor_to_graph<'a>(
142    _ctx: &'a AtomArena<'a>,
143    expr: Atom<'a>,
144    registry: &TensorRegistry,
145) -> Result<
146    (
147        Graph<TgNode, usize, TgEdge>,
148        Vec<HeadInfo>,
149        HashMap<usize, Atom<'a>>,
150    ),
151    TensorCanonError,
152> {
153    let mut g: Graph<TgNode, usize, TgEdge> = Graph::new();
154    let mut heads: Vec<HeadInfo> = Vec::new();
155    let mut index_uses: HashMap<Atom<'a>, (Vec<usize>, usize)> = HashMap::new();
156    let mut slot_labels: HashMap<usize, Atom<'a>> = HashMap::new();
157
158    match expr.node() {
159        AtomNode::Mul(factors) => {
160            for f in factors.iter() {
161                encode_factor(
162                    *f,
163                    registry,
164                    &mut g,
165                    &mut heads,
166                    &mut index_uses,
167                    &mut slot_labels,
168                )?;
169            }
170        }
171        _ => {
172            encode_factor(
173                expr,
174                registry,
175                &mut g,
176                &mut heads,
177                &mut index_uses,
178                &mut slot_labels,
179            )?;
180        }
181    }
182
183    for (_label, (slot_verts, count)) in &index_uses {
184        if *count > 2 {
185            return Err(TensorCanonError::ContractedMoreThanOnce(Symbol::new(
186                &_label.to_string(),
187            )));
188        }
189        if *count == 2 {
190            let group = registry.index_group(Symbol::new(&_label.to_string()));
191            g.add_undirected_edge(slot_verts[0], slot_verts[1], TgEdge::Contraction(group));
192        }
193    }
194
195    Ok((g, heads, slot_labels))
196}
197
198fn encode_factor<'a>(
199    factor: Atom<'a>,
200    registry: &TensorRegistry,
201    g: &mut Graph<TgNode, usize, TgEdge>,
202    heads: &mut Vec<HeadInfo>,
203    index_uses: &mut HashMap<Atom<'a>, (Vec<usize>, usize)>,
204    slot_labels: &mut HashMap<usize, Atom<'a>>,
205) -> Result<(), TensorCanonError> {
206    match factor.node() {
207        AtomNode::Fun(name, args) => {
208            let spec = registry
209                .spec(*name)
210                .ok_or(TensorCanonError::NotATensor(*name))?;
211            let head_v = g.add_node(TgNode::Head(hash(name.as_str())), 0);
212            let mut slot_verts = Vec::with_capacity(args.len());
213
214            // Pre-sort symmetric slots by label so the graph encoding is
215            // input-order-independent.  Non-symmetric slots keep their
216            // original position.
217            let mut sorted_args: Vec<(usize, Atom<'a>)> =
218                args.iter().enumerate().map(|(i, a)| (i, *a)).collect();
219            // Stable sort: symmetric slots sorted by label, others by pos.
220            sorted_args.sort_by(|&(pa, aa), &(pb, ab)| {
221                let ha = spec.is_slot_hidden(pa);
222                let hb = spec.is_slot_hidden(pb);
223                match (ha, hb) {
224                    (true, true) => {
225                        // Compare by Symbol name to avoid platform-dependent
226                        // Display differences.
227                        let sa = match aa.node() {
228                            AtomNode::Var(s) => s.as_str(),
229                            _ => "",
230                        };
231                        let sb = match ab.node() {
232                            AtomNode::Var(s) => s.as_str(),
233                            _ => "",
234                        };
235                        sa.cmp(sb)
236                    }
237                    (true, false) => std::cmp::Ordering::Less,
238                    (false, true) => std::cmp::Ordering::Greater,
239                    (false, false) => pa.cmp(&pb),
240                }
241            });
242
243            for (sorted_idx, (orig_pos, arg)) in sorted_args.into_iter().enumerate() {
244                let label = arg;
245                let is_hidden = spec.is_slot_hidden(orig_pos);
246                // Symmetric slots must share the same colour so the graph-
247                // isomorphism engine can freely permute them, ensuring
248                // canonicalise(g(a,b)) == canonicalise(g(b,a)).
249                let slot_colour = if is_hidden {
250                    TgNode::Slot(0)
251                } else {
252                    TgNode::Slot(hash(&label.to_string()))
253                };
254                let slot_v = g.add_node(slot_colour, 0);
255                slot_labels.insert(slot_v, label);
256                slot_verts.push(slot_v);
257
258                // Use sorted index as edge pos so the graph encoding is
259                // identical for symmetric inputs like g(a,b) and g(b,a).
260                let edge_pos = if is_hidden { sorted_idx } else { orig_pos };
261                let kind = if is_hidden {
262                    TgEdge::HeadToSlot(edge_pos, 0)
263                } else {
264                    TgEdge::HeadToSlot(edge_pos, 1)
265                };
266                g.add_directed_edge(head_v, slot_v, kind);
267
268                let entry = index_uses.entry(label).or_insert_with(|| (Vec::new(), 0));
269                entry.0.push(slot_v);
270                entry.1 += 1;
271            }
272
273            heads.push(HeadInfo {
274                symbol: *name,
275                slot_count: args.len(),
276                head_v,
277                slot_verts,
278            });
279        }
280        AtomNode::Pow(_, _) => return Err(TensorCanonError::UnsupportedPower),
281        _ => {
282            let h = hash(&factor.to_string());
283            g.add_node(TgNode::Scalar(h), 0);
284        }
285    }
286    Ok(())
287}
288
289// =========================================================================
290// Reconstruction
291// =========================================================================
292
293#[allow(clippy::type_complexity, clippy::needless_range_loop)]
294fn reconstruct<'a>(
295    ctx: &'a AtomArena<'a>,
296    cf: &CanonicalForm<TgNode, usize, TgEdge>,
297    heads: &[HeadInfo],
298    slot_labels: &HashMap<usize, Atom<'a>>,
299    _registry: &TensorRegistry,
300) -> Result<CanonicalTensor<'a>, TensorCanonError> {
301    let cg = &cf.graph;
302    let n = cg.node_count();
303
304    // Map canonical → original vertex (vertex_map[pos] = original_vertex).
305    let orig_of = &cf.vertex_map;
306
307    // Find head→slot edges and contraction pairs in canonical graph.
308    let mut slot_contraction: HashMap<usize, (usize, u64)> = HashMap::new();
309    for v in 0..n {
310        for ev in cg.edges_of(v) {
311            if !ev.is_directed
312                && let TgEdge::Contraction(g) = ev.data
313            {
314                slot_contraction.insert(v, (ev.neighbour, g));
315            }
316        }
317    }
318
319    // Assign canonical dummy names to each contraction pair.
320    let mut group_counters: HashMap<u64, usize> = HashMap::new();
321    // Key: (min(cv1, cv2), max(cv1, cv2))
322    let mut pair_labels: HashMap<(usize, usize), Atom<'a>> = HashMap::new();
323
324    for v in 0..n {
325        for ev in cg.edges_of(v) {
326            if !ev.is_directed
327                && let TgEdge::Contraction(g) = ev.data
328            {
329                let a = v.min(ev.neighbour);
330                let b = v.max(ev.neighbour);
331                pair_labels.entry((a, b)).or_insert_with(|| {
332                    let cnt = group_counters.entry(g).or_insert(0);
333                    let label = if g == 0 {
334                        ctx.var(&format!("d{}", cnt))
335                    } else {
336                        ctx.var(&format!("d{}_{}", g, cnt))
337                    };
338                    *cnt += 1;
339                    label
340                });
341            }
342        }
343    }
344
345    // Collect canonical heads and their slots.
346    let mut canon_heads: Vec<(usize, &HeadInfo)> = Vec::new();
347    let mut orig_to_head: HashMap<usize, &HeadInfo> = HashMap::new();
348    for h in heads {
349        orig_to_head.insert(h.head_v, h);
350    }
351    for v in 0..n {
352        if let TgNode::Head(_) = cg.node_data(v) {
353            let orig = orig_of[v];
354            if let Some(h) = orig_to_head.get(&orig) {
355                canon_heads.push((v, *h));
356            }
357        }
358    }
359    canon_heads.sort_by_key(|(v, _)| *v);
360
361    // Build factors.
362    let mut factors: Vec<Atom<'a>> = Vec::new();
363    let mut all_dummies: Vec<Atom<'a>> = Vec::new();
364    let mut external_indices: Vec<Atom<'a>> = Vec::new();
365
366    for (can_head, h) in &canon_heads {
367        // Gather slot info.
368        let mut slot_infos: Vec<SlotInfo> = Vec::new();
369        for ev in cg.edges_of(*can_head) {
370            if ev.is_directed
371                && ev.is_outgoing
372                && let TgEdge::HeadToSlot(pos, hidden_flag) = ev.data
373            {
374                let partner = slot_contraction.get(&ev.neighbour).copied();
375                slot_infos.push(SlotInfo {
376                    orig_pos: pos,
377                    hidden: hidden_flag == 0,
378                    canon_slot_v: ev.neighbour,
379                    partner_v: partner.map(|(p, _)| p),
380                });
381            }
382        }
383
384        // Sort: hidden (symmetric) slots by label string (deterministic
385        // regardless of input order), visible slots by original position.
386        slot_infos.sort_by(|a, b| match (a.hidden, b.hidden) {
387            (true, true) => {
388                let la = orig_of[a.canon_slot_v];
389                let lb = orig_of[b.canon_slot_v];
390                let sa = slot_labels
391                    .get(&la)
392                    .map(|x| x.to_string())
393                    .unwrap_or_default();
394                let sb = slot_labels
395                    .get(&lb)
396                    .map(|x| x.to_string())
397                    .unwrap_or_default();
398                sa.cmp(&sb)
399            }
400            (true, false) => std::cmp::Ordering::Less,
401            (false, true) => std::cmp::Ordering::Greater,
402            (false, false) => a.orig_pos.cmp(&b.orig_pos),
403        });
404
405        let mut args: Vec<Atom<'a>> = Vec::new();
406        for si in &slot_infos {
407            if let Some(pv) = si.partner_v {
408                let a = si.canon_slot_v.min(pv);
409                let b = si.canon_slot_v.max(pv);
410                if let Some(label) = pair_labels.get(&(a, b)) {
411                    args.push(*label);
412                    if !all_dummies.contains(label) {
413                        all_dummies.push(*label);
414                    }
415                } else {
416                    args.push(ctx.var("?"));
417                }
418            } else {
419                // External index: preserve original label from the graph encoding.
420                let orig_slot = orig_of[si.canon_slot_v];
421                if let Some(&orig_label) = slot_labels.get(&orig_slot) {
422                    args.push(orig_label);
423                    if !external_indices.contains(&orig_label) {
424                        external_indices.push(orig_label);
425                    }
426                } else {
427                    // Fallback: synthetic name.
428                    let label = ctx.var(&format!("ext{}", external_indices.len()));
429                    args.push(label);
430                    if !external_indices.contains(&label) {
431                        external_indices.push(label);
432                    }
433                }
434            }
435        }
436
437        factors.push(ctx.fun(h.symbol.as_str(), &args));
438    }
439
440    let canonical_form = if factors.is_empty() {
441        ctx.num(1)
442    } else if factors.len() == 1 {
443        factors.pop().unwrap()
444    } else {
445        ctx.mul(&factors)
446    };
447
448    Ok(CanonicalTensor {
449        canonical_form,
450        external_indices,
451        dummy_indices: all_dummies,
452    })
453}
454
455struct SlotInfo {
456    orig_pos: usize,
457    hidden: bool,
458    canon_slot_v: usize,
459    partner_v: Option<usize>,
460}
461
462fn hash(s: &str) -> u64 {
463    use std::hash::Hasher;
464    let mut h = std::collections::hash_map::DefaultHasher::new();
465    std::hash::Hash::hash(&s, &mut h);
466    h.finish()
467}
468
469// =========================================================================
470// Tests
471// =========================================================================
472
473#[cfg(test)]
474mod tests {
475    use super::*;
476    use crate::AtomArena;
477    use crate::Symbol;
478    use crate::tensor::spec::SymmetrySpec;
479    use ocas_core::arena::Arena;
480
481    #[test]
482    fn canon_single_tensor_no_symmetry() {
483        let arena = Arena::new();
484        let ctx = AtomArena::new(&arena);
485        let mut reg = TensorRegistry::new();
486        reg.register(Symbol::new("T"), SymmetrySpec::none());
487
488        let i = ctx.var("i");
489        let j = ctx.var("j");
490        let t = ctx.fun("T", &[i, j]);
491        let ct = canonicalize_tensors(&ctx, t, &reg).unwrap();
492        let s = ct.canonical_form.to_string();
493        assert!(s.contains("T"), "result: {s}");
494    }
495
496    #[test]
497    fn canon_product_with_contraction() {
498        let arena = Arena::new();
499        let ctx = AtomArena::new(&arena);
500        let mut reg = TensorRegistry::new();
501        reg.register(Symbol::new("T"), SymmetrySpec::none());
502        reg.register(Symbol::new("U"), SymmetrySpec::none());
503
504        let i = ctx.var("i");
505        let j = ctx.var("j");
506        let k = ctx.var("k");
507        let t = ctx.fun("T", &[i, j]);
508        let u = ctx.fun("U", &[j, k]);
509        let prod = ctx.mul(&[t, u]);
510        let ct = canonicalize_tensors(&ctx, prod, &reg).unwrap();
511        let s = ct.canonical_form.to_string();
512        // Should have at least one dummy and two tensors.
513        assert!(s.contains("d0"), "expected dummy d0, got: {s}");
514        assert!(s.contains('T') && s.contains('U'), "got: {s}");
515        assert_eq!(ct.dummy_indices.len(), 1);
516    }
517
518    #[test]
519    fn canon_symmetric_tensor_consistency() {
520        let arena = Arena::new();
521        let ctx = AtomArena::new(&arena);
522        let mut reg = TensorRegistry::new();
523        reg.register(Symbol::new("g"), SymmetrySpec::fully_symmetric(2));
524
525        let a = ctx.var("a");
526        let b = ctx.var("b");
527        let g_ab = ctx.fun("g", &[a, b]);
528        let g_ba = ctx.fun("g", &[b, a]);
529        let ct1 = canonicalize_tensors(&ctx, g_ab, &reg).unwrap();
530        let ct2 = canonicalize_tensors(&ctx, g_ba, &reg).unwrap();
531        // Both should canonicalise to the same form.
532        assert_eq!(
533            ct1.canonical_form.to_string(),
534            ct2.canonical_form.to_string(),
535            "symmetric slots should canonicalise consistently"
536        );
537    }
538}