use std::collections::HashMap;
use crate::{Atom, AtomArena, AtomNode, Symbol};
use super::graph::{CanonicalForm, Graph};
use super::spec::TensorRegistry;
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TensorCanonError {
ContractedMoreThanOnce(Symbol),
BadContraction(Symbol),
NotATensor(Symbol),
InconsistentOpenIndices,
UnsupportedPower,
}
#[derive(Debug, Clone)]
pub struct CanonicalTensor<'a> {
pub canonical_form: Atom<'a>,
pub external_indices: Vec<Atom<'a>>,
pub dummy_indices: Vec<Atom<'a>>,
}
pub fn canonicalize_tensors<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
registry: &TensorRegistry,
) -> Result<CanonicalTensor<'a>, TensorCanonError> {
match expr.node() {
AtomNode::Add(terms) => {
let mut canon_terms: Vec<Atom<'a>> = Vec::new();
let mut first_external: Option<Vec<Atom<'a>>> = None;
let mut all_dummies: Vec<Atom<'a>> = Vec::new();
for term in terms.iter() {
let ct = canonicalize_single_term(ctx, *term, registry)?;
match &first_external {
None => first_external = Some(ct.external_indices.clone()),
Some(ext) if *ext != ct.external_indices => {
return Err(TensorCanonError::InconsistentOpenIndices);
}
_ => {}
}
all_dummies.extend(ct.dummy_indices);
canon_terms.push(ct.canonical_form);
}
let canonical_form = if canon_terms.len() == 1 {
canon_terms.pop().unwrap()
} else {
ctx.add(&canon_terms)
};
Ok(CanonicalTensor {
canonical_form,
external_indices: first_external.unwrap_or_default(),
dummy_indices: all_dummies,
})
}
_ => canonicalize_single_term(ctx, expr, registry),
}
}
fn canonicalize_single_term<'a>(
ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
registry: &TensorRegistry,
) -> Result<CanonicalTensor<'a>, TensorCanonError> {
#[allow(clippy::collapsible_if)]
if let AtomNode::Fun(name, args) = expr.node() {
if let Some(spec) = registry.spec(*name) {
let all_symmetric = !spec.symmetric_subsets.is_empty()
&& spec.antisymmetric_subsets.is_empty()
&& (0..args.len()).all(|pos| spec.is_slot_hidden(pos));
if all_symmetric {
let mut sorted: Vec<Atom<'a>> = args.to_vec();
sorted.sort_by_key(|a| match a.node() {
AtomNode::Var(s) => s.as_str().to_string(),
_ => a.to_string(),
});
let result = ctx.fun(name.as_str(), &sorted);
return Ok(CanonicalTensor {
canonical_form: result,
external_indices: sorted,
dummy_indices: Vec::new(),
});
}
}
}
let (g, head_nodes, slot_labels) = tensor_to_graph(ctx, expr, registry)?;
let cf = g.canonize();
reconstruct(ctx, &cf, &head_nodes, &slot_labels, registry)
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
enum TgNode {
Head(u64),
Slot(u64),
Scalar(u64),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
enum TgEdge {
HeadToSlot(usize, u8),
Contraction(u64),
}
#[derive(Debug, Clone)]
#[allow(dead_code)]
struct HeadInfo {
symbol: Symbol,
slot_count: usize,
head_v: usize,
slot_verts: Vec<usize>,
}
#[allow(clippy::type_complexity)]
fn tensor_to_graph<'a>(
_ctx: &'a AtomArena<'a>,
expr: Atom<'a>,
registry: &TensorRegistry,
) -> Result<
(
Graph<TgNode, usize, TgEdge>,
Vec<HeadInfo>,
HashMap<usize, Atom<'a>>,
),
TensorCanonError,
> {
let mut g: Graph<TgNode, usize, TgEdge> = Graph::new();
let mut heads: Vec<HeadInfo> = Vec::new();
let mut index_uses: HashMap<Atom<'a>, (Vec<usize>, usize)> = HashMap::new();
let mut slot_labels: HashMap<usize, Atom<'a>> = HashMap::new();
match expr.node() {
AtomNode::Mul(factors) => {
for f in factors.iter() {
encode_factor(
*f,
registry,
&mut g,
&mut heads,
&mut index_uses,
&mut slot_labels,
)?;
}
}
_ => {
encode_factor(
expr,
registry,
&mut g,
&mut heads,
&mut index_uses,
&mut slot_labels,
)?;
}
}
for (_label, (slot_verts, count)) in &index_uses {
if *count > 2 {
return Err(TensorCanonError::ContractedMoreThanOnce(Symbol::new(
&_label.to_string(),
)));
}
if *count == 2 {
let group = registry.index_group(Symbol::new(&_label.to_string()));
g.add_undirected_edge(slot_verts[0], slot_verts[1], TgEdge::Contraction(group));
}
}
Ok((g, heads, slot_labels))
}
fn encode_factor<'a>(
factor: Atom<'a>,
registry: &TensorRegistry,
g: &mut Graph<TgNode, usize, TgEdge>,
heads: &mut Vec<HeadInfo>,
index_uses: &mut HashMap<Atom<'a>, (Vec<usize>, usize)>,
slot_labels: &mut HashMap<usize, Atom<'a>>,
) -> Result<(), TensorCanonError> {
match factor.node() {
AtomNode::Fun(name, args) => {
let spec = registry
.spec(*name)
.ok_or(TensorCanonError::NotATensor(*name))?;
let head_v = g.add_node(TgNode::Head(hash(name.as_str())), 0);
let mut slot_verts = Vec::with_capacity(args.len());
let mut sorted_args: Vec<(usize, Atom<'a>)> =
args.iter().enumerate().map(|(i, a)| (i, *a)).collect();
sorted_args.sort_by(|&(pa, aa), &(pb, ab)| {
let ha = spec.is_slot_hidden(pa);
let hb = spec.is_slot_hidden(pb);
match (ha, hb) {
(true, true) => {
let sa = match aa.node() {
AtomNode::Var(s) => s.as_str(),
_ => "",
};
let sb = match ab.node() {
AtomNode::Var(s) => s.as_str(),
_ => "",
};
sa.cmp(sb)
}
(true, false) => std::cmp::Ordering::Less,
(false, true) => std::cmp::Ordering::Greater,
(false, false) => pa.cmp(&pb),
}
});
for (sorted_idx, (orig_pos, arg)) in sorted_args.into_iter().enumerate() {
let label = arg;
let is_hidden = spec.is_slot_hidden(orig_pos);
let slot_colour = if is_hidden {
TgNode::Slot(0)
} else {
TgNode::Slot(hash(&label.to_string()))
};
let slot_v = g.add_node(slot_colour, 0);
slot_labels.insert(slot_v, label);
slot_verts.push(slot_v);
let edge_pos = if is_hidden { sorted_idx } else { orig_pos };
let kind = if is_hidden {
TgEdge::HeadToSlot(edge_pos, 0)
} else {
TgEdge::HeadToSlot(edge_pos, 1)
};
g.add_directed_edge(head_v, slot_v, kind);
let entry = index_uses.entry(label).or_insert_with(|| (Vec::new(), 0));
entry.0.push(slot_v);
entry.1 += 1;
}
heads.push(HeadInfo {
symbol: *name,
slot_count: args.len(),
head_v,
slot_verts,
});
}
AtomNode::Pow(_, _) => return Err(TensorCanonError::UnsupportedPower),
_ => {
let h = hash(&factor.to_string());
g.add_node(TgNode::Scalar(h), 0);
}
}
Ok(())
}
#[allow(clippy::type_complexity, clippy::needless_range_loop)]
fn reconstruct<'a>(
ctx: &'a AtomArena<'a>,
cf: &CanonicalForm<TgNode, usize, TgEdge>,
heads: &[HeadInfo],
slot_labels: &HashMap<usize, Atom<'a>>,
_registry: &TensorRegistry,
) -> Result<CanonicalTensor<'a>, TensorCanonError> {
let cg = &cf.graph;
let n = cg.node_count();
let orig_of = &cf.vertex_map;
let mut slot_contraction: HashMap<usize, (usize, u64)> = HashMap::new();
for v in 0..n {
for ev in cg.edges_of(v) {
if !ev.is_directed
&& let TgEdge::Contraction(g) = ev.data
{
slot_contraction.insert(v, (ev.neighbour, g));
}
}
}
let mut group_counters: HashMap<u64, usize> = HashMap::new();
let mut pair_labels: HashMap<(usize, usize), Atom<'a>> = HashMap::new();
for v in 0..n {
for ev in cg.edges_of(v) {
if !ev.is_directed
&& let TgEdge::Contraction(g) = ev.data
{
let a = v.min(ev.neighbour);
let b = v.max(ev.neighbour);
pair_labels.entry((a, b)).or_insert_with(|| {
let cnt = group_counters.entry(g).or_insert(0);
let label = if g == 0 {
ctx.var(&format!("d{}", cnt))
} else {
ctx.var(&format!("d{}_{}", g, cnt))
};
*cnt += 1;
label
});
}
}
}
let mut canon_heads: Vec<(usize, &HeadInfo)> = Vec::new();
let mut orig_to_head: HashMap<usize, &HeadInfo> = HashMap::new();
for h in heads {
orig_to_head.insert(h.head_v, h);
}
for v in 0..n {
if let TgNode::Head(_) = cg.node_data(v) {
let orig = orig_of[v];
if let Some(h) = orig_to_head.get(&orig) {
canon_heads.push((v, *h));
}
}
}
canon_heads.sort_by_key(|(v, _)| *v);
let mut factors: Vec<Atom<'a>> = Vec::new();
let mut all_dummies: Vec<Atom<'a>> = Vec::new();
let mut external_indices: Vec<Atom<'a>> = Vec::new();
for (can_head, h) in &canon_heads {
let mut slot_infos: Vec<SlotInfo> = Vec::new();
for ev in cg.edges_of(*can_head) {
if ev.is_directed
&& ev.is_outgoing
&& let TgEdge::HeadToSlot(pos, hidden_flag) = ev.data
{
let partner = slot_contraction.get(&ev.neighbour).copied();
slot_infos.push(SlotInfo {
orig_pos: pos,
hidden: hidden_flag == 0,
canon_slot_v: ev.neighbour,
partner_v: partner.map(|(p, _)| p),
});
}
}
slot_infos.sort_by(|a, b| match (a.hidden, b.hidden) {
(true, true) => {
let la = orig_of[a.canon_slot_v];
let lb = orig_of[b.canon_slot_v];
let sa = slot_labels
.get(&la)
.map(|x| x.to_string())
.unwrap_or_default();
let sb = slot_labels
.get(&lb)
.map(|x| x.to_string())
.unwrap_or_default();
sa.cmp(&sb)
}
(true, false) => std::cmp::Ordering::Less,
(false, true) => std::cmp::Ordering::Greater,
(false, false) => a.orig_pos.cmp(&b.orig_pos),
});
let mut args: Vec<Atom<'a>> = Vec::new();
for si in &slot_infos {
if let Some(pv) = si.partner_v {
let a = si.canon_slot_v.min(pv);
let b = si.canon_slot_v.max(pv);
if let Some(label) = pair_labels.get(&(a, b)) {
args.push(*label);
if !all_dummies.contains(label) {
all_dummies.push(*label);
}
} else {
args.push(ctx.var("?"));
}
} else {
let orig_slot = orig_of[si.canon_slot_v];
if let Some(&orig_label) = slot_labels.get(&orig_slot) {
args.push(orig_label);
if !external_indices.contains(&orig_label) {
external_indices.push(orig_label);
}
} else {
let label = ctx.var(&format!("ext{}", external_indices.len()));
args.push(label);
if !external_indices.contains(&label) {
external_indices.push(label);
}
}
}
}
factors.push(ctx.fun(h.symbol.as_str(), &args));
}
let canonical_form = if factors.is_empty() {
ctx.num(1)
} else if factors.len() == 1 {
factors.pop().unwrap()
} else {
ctx.mul(&factors)
};
Ok(CanonicalTensor {
canonical_form,
external_indices,
dummy_indices: all_dummies,
})
}
struct SlotInfo {
orig_pos: usize,
hidden: bool,
canon_slot_v: usize,
partner_v: Option<usize>,
}
fn hash(s: &str) -> u64 {
use std::hash::Hasher;
let mut h = std::collections::hash_map::DefaultHasher::new();
std::hash::Hash::hash(&s, &mut h);
h.finish()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::AtomArena;
use crate::Symbol;
use crate::tensor::spec::SymmetrySpec;
use ocas_core::arena::Arena;
#[test]
fn canon_single_tensor_no_symmetry() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let mut reg = TensorRegistry::new();
reg.register(Symbol::new("T"), SymmetrySpec::none());
let i = ctx.var("i");
let j = ctx.var("j");
let t = ctx.fun("T", &[i, j]);
let ct = canonicalize_tensors(&ctx, t, ®).unwrap();
let s = ct.canonical_form.to_string();
assert!(s.contains("T"), "result: {s}");
}
#[test]
fn canon_product_with_contraction() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let mut reg = TensorRegistry::new();
reg.register(Symbol::new("T"), SymmetrySpec::none());
reg.register(Symbol::new("U"), SymmetrySpec::none());
let i = ctx.var("i");
let j = ctx.var("j");
let k = ctx.var("k");
let t = ctx.fun("T", &[i, j]);
let u = ctx.fun("U", &[j, k]);
let prod = ctx.mul(&[t, u]);
let ct = canonicalize_tensors(&ctx, prod, ®).unwrap();
let s = ct.canonical_form.to_string();
assert!(s.contains("d0"), "expected dummy d0, got: {s}");
assert!(s.contains('T') && s.contains('U'), "got: {s}");
assert_eq!(ct.dummy_indices.len(), 1);
}
#[test]
fn canon_symmetric_tensor_consistency() {
let arena = Arena::new();
let ctx = AtomArena::new(&arena);
let mut reg = TensorRegistry::new();
reg.register(Symbol::new("g"), SymmetrySpec::fully_symmetric(2));
let a = ctx.var("a");
let b = ctx.var("b");
let g_ab = ctx.fun("g", &[a, b]);
let g_ba = ctx.fun("g", &[b, a]);
let ct1 = canonicalize_tensors(&ctx, g_ab, ®).unwrap();
let ct2 = canonicalize_tensors(&ctx, g_ba, ®).unwrap();
assert_eq!(
ct1.canonical_form.to_string(),
ct2.canonical_form.to_string(),
"symmetric slots should canonicalise consistently"
);
}
}