1use std::collections::HashMap;
11
12use crate::{Atom, AtomArena, AtomNode, Symbol};
13
14use super::graph::{CanonicalForm, Graph};
15use super::spec::TensorRegistry;
16
17#[derive(Debug, Clone, PartialEq, Eq)]
19pub enum TensorCanonError {
20 ContractedMoreThanOnce(Symbol),
21 BadContraction(Symbol),
22 NotATensor(Symbol),
23 InconsistentOpenIndices,
24 UnsupportedPower,
25}
26
27#[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
35pub 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 #[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#[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#[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 let mut sorted_args: Vec<(usize, Atom<'a>)> =
218 args.iter().enumerate().map(|(i, a)| (i, *a)).collect();
219 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 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 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 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#[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 let orig_of = &cf.vertex_map;
306
307 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 let mut group_counters: HashMap<u64, usize> = HashMap::new();
321 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 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 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 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 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 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 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#[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, ®).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, ®).unwrap();
511 let s = ct.canonical_form.to_string();
512 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, ®).unwrap();
530 let ct2 = canonicalize_tensors(&ctx, g_ba, ®).unwrap();
531 assert_eq!(
533 ct1.canonical_form.to_string(),
534 ct2.canonical_form.to_string(),
535 "symmetric slots should canonicalise consistently"
536 );
537 }
538}