pub mod canon;
pub mod dummy;
pub mod graph;
pub mod spec;
pub mod young;
use crate::{Atom, AtomArena, Symbol};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum IndexPosition {
Upper,
Lower,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct IndexSlot<'a> {
label: Atom<'a>,
position: IndexPosition,
}
impl<'a> IndexSlot<'a> {
pub fn new(label: Atom<'a>, position: IndexPosition) -> Self {
Self { label, position }
}
pub fn label(&self) -> Atom<'a> {
self.label
}
pub fn position(&self) -> IndexPosition {
self.position
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Symmetry {
None,
Symmetric,
Antisymmetric,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Tensor<'a> {
name: Symbol,
slots: Vec<IndexSlot<'a>>,
symmetry: Symmetry,
}
impl<'a> Tensor<'a> {
pub fn new(name: Symbol, slots: Vec<IndexSlot<'a>>) -> Self {
Self {
name,
slots,
symmetry: Symmetry::None,
}
}
pub fn with_symmetry(mut self, symmetry: Symmetry) -> Self {
self.symmetry = symmetry;
self
}
pub fn name(&self) -> Symbol {
self.name
}
pub fn slots(&self) -> &[IndexSlot<'a>] {
&self.slots
}
pub fn symmetry(&self) -> Symmetry {
self.symmetry
}
pub fn rank(&self) -> usize {
self.slots.len()
}
pub fn dummy_labels(&self) -> Vec<Atom<'a>> {
dummies(self.slots().iter().map(|s| s.label()))
}
pub fn to_atom(&self, ctx: &'a AtomArena<'a>) -> Atom<'a> {
let args: Vec<Atom<'a>> = self.slots.iter().map(|s| s.label).collect();
ctx.fun(self.name.as_str(), &args)
}
}
fn dummies<'a, I: IntoIterator<Item = Atom<'a>>>(labels: I) -> Vec<Atom<'a>> {
use crate::FastHashMap;
let mut counts: FastHashMap<AtomId<'a>, usize> = FastHashMap::default();
for l in labels {
let id = AtomId(l);
*counts.entry(id).or_insert(0) += 1;
}
let mut out: Vec<Atom<'a>> = Vec::new();
let mut seen: std::collections::HashSet<*const ()> = std::collections::HashSet::new();
for (id, n) in counts.iter() {
if *n == 2 {
let ptr = id.0.node() as *const _ as *const ();
if seen.insert(ptr) {
out.push(id.0);
}
}
}
out
}
#[derive(Clone, Copy, PartialEq, Eq, Hash)]
struct AtomId<'a>(Atom<'a>);
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum Contracted<'a> {
Product(TensorProduct<'a>),
Scalar(Atom<'a>),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TensorProduct<'a> {
pub factors: Vec<Tensor<'a>>,
}
pub fn contract<'a>(ctx: &'a AtomArena<'a>, a: &Tensor<'a>, b: &Tensor<'a>) -> Contracted<'a> {
let mut used_a = vec![false; a.slots.len()];
let mut used_b = vec![false; b.slots.len()];
let mut pair_labels: Vec<Atom<'a>> = Vec::new();
for (i, sa) in a.slots.iter().enumerate() {
if used_a[i] {
continue;
}
for (j, sb) in b.slots.iter().enumerate() {
if used_b[j] {
continue;
}
if sa.label == sb.label && sa.position != sb.position {
used_a[i] = true;
used_b[j] = true;
pair_labels.push(sa.label);
break;
}
}
}
let mut free: Vec<IndexSlot<'a>> = Vec::new();
for (i, s) in a.slots.iter().enumerate() {
if !used_a[i] {
free.push(*s);
}
}
for (j, s) in b.slots.iter().enumerate() {
if !used_b[j] {
free.push(*s);
}
}
if pair_labels.is_empty() {
return Contracted::Product(TensorProduct {
factors: vec![a.clone(), b.clone()],
});
}
if free.is_empty() {
let a_atom = a.to_atom(ctx);
let b_atom = b.to_atom(ctx);
let product = ctx.mul(&[a_atom, b_atom]);
return Contracted::Scalar(product);
}
let name = Symbol::new(&format!("{}_contract_{}", a.name.as_str(), b.name.as_str()));
Contracted::Product(TensorProduct {
factors: vec![Tensor::new(name, free)],
})
}
pub fn symmetrise_sign(tensor: &Tensor<'_>) -> i64 {
match tensor.symmetry {
Symmetry::None | Symmetry::Symmetric => 1,
Symmetry::Antisymmetric => {
let mut slots: Vec<IndexSlot<'_>> = tensor.slots.to_vec();
let mut swaps = 0usize;
for i in 1..slots.len() {
let mut j = i;
while j > 0 && slot_less(&slots[j - 1], &slots[j]) {
slots.swap(j - 1, j);
swaps += 1;
j -= 1;
}
}
if swaps.is_multiple_of(2) { 1 } else { -1 }
}
}
}
fn slot_less(a: &IndexSlot<'_>, b: &IndexSlot<'_>) -> bool {
let pa = a.label.node() as *const _ as *const ();
let pb = b.label.node() as *const _ as *const ();
(pa as usize) < (pb as usize) || (pa == pb && (a.position as u8) > (b.position as u8))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::AtomArena;
use crate::AtomNode;
fn idx<'a>(ctx: &'a AtomArena<'a>, name: &str, pos: IndexPosition) -> IndexSlot<'a> {
IndexSlot::new(ctx.var(name), pos)
}
#[test]
fn tensor_rank_and_slots() {
let arena = crate::Arena::new();
let ctx = AtomArena::new(&arena);
let t = Tensor::new(
Symbol::new("T"),
vec![
idx(&ctx, "i", IndexPosition::Upper),
idx(&ctx, "j", IndexPosition::Lower),
],
);
assert_eq!(t.rank(), 2);
assert_eq!(t.slots().len(), 2);
assert_eq!(t.symmetry(), Symmetry::None);
}
#[test]
fn dummy_detection_finds_repeated_label() {
let arena = crate::Arena::new();
let ctx = AtomArena::new(&arena);
let t = Tensor::new(
Symbol::new("T"),
vec![
idx(&ctx, "i", IndexPosition::Upper),
idx(&ctx, "i", IndexPosition::Lower),
],
);
let dummies = t.dummy_labels();
assert_eq!(dummies.len(), 1);
}
#[test]
fn contract_two_tensors_with_one_dummy() {
let arena = crate::Arena::new();
let ctx = AtomArena::new(&arena);
let t = Tensor::new(
Symbol::new("T"),
vec![
idx(&ctx, "i", IndexPosition::Upper),
idx(&ctx, "j", IndexPosition::Lower),
],
);
let u = Tensor::new(
Symbol::new("U"),
vec![
idx(&ctx, "j", IndexPosition::Upper),
idx(&ctx, "k", IndexPosition::Lower),
],
);
match contract(&ctx, &t, &u) {
Contracted::Product(p) => {
assert_eq!(p.factors.len(), 1);
assert_eq!(p.factors[0].rank(), 2);
}
_ => panic!("expected partial contraction product"),
}
}
#[test]
fn contract_to_scalar_when_no_free_slots() {
let arena = crate::Arena::new();
let ctx = AtomArena::new(&arena);
let t = Tensor::new(Symbol::new("T"), vec![idx(&ctx, "i", IndexPosition::Upper)]);
let u = Tensor::new(Symbol::new("U"), vec![idx(&ctx, "i", IndexPosition::Lower)]);
match contract(&ctx, &t, &u) {
Contracted::Scalar(atom) => {
assert!(matches!(atom.node(), AtomNode::Mul(_)));
}
_ => panic!("expected scalar contraction"),
}
}
#[test]
fn no_overlap_yields_plain_product() {
let arena = crate::Arena::new();
let ctx = AtomArena::new(&arena);
let t = Tensor::new(Symbol::new("T"), vec![idx(&ctx, "i", IndexPosition::Upper)]);
let u = Tensor::new(Symbol::new("U"), vec![idx(&ctx, "j", IndexPosition::Upper)]);
match contract(&ctx, &t, &u) {
Contracted::Product(p) => assert_eq!(p.factors.len(), 2),
_ => panic!("expected plain product"),
}
}
#[test]
fn antisymmetric_sign_parity() {
let arena = crate::Arena::new();
let ctx = AtomArena::new(&arena);
let e_ab = Tensor::new(
Symbol::new("eps"),
vec![
idx(&ctx, "a", IndexPosition::Lower),
idx(&ctx, "b", IndexPosition::Lower),
],
)
.with_symmetry(Symmetry::Antisymmetric);
let e_ba = Tensor::new(
Symbol::new("eps"),
vec![
idx(&ctx, "b", IndexPosition::Lower),
idx(&ctx, "a", IndexPosition::Lower),
],
)
.with_symmetry(Symmetry::Antisymmetric);
let s1 = symmetrise_sign(&e_ab);
let s2 = symmetrise_sign(&e_ba);
assert!(s1 == 1 || s1 == -1);
assert!(s2 == 1 || s2 == -1);
assert_eq!(s1, -s2);
}
#[test]
fn symmetric_sign_is_always_plus() {
let arena = crate::Arena::new();
let ctx = AtomArena::new(&arena);
let g = Tensor::new(
Symbol::new("g"),
vec![
idx(&ctx, "a", IndexPosition::Lower),
idx(&ctx, "b", IndexPosition::Lower),
],
)
.with_symmetry(Symmetry::Symmetric);
assert_eq!(symmetrise_sign(&g), 1);
}
#[test]
fn to_atom_round_trips_as_function_node() {
let arena = crate::Arena::new();
let ctx = AtomArena::new(&arena);
let t = Tensor::new(
Symbol::new("T"),
vec![
idx(&ctx, "i", IndexPosition::Upper),
idx(&ctx, "j", IndexPosition::Lower),
],
);
let atom = t.to_atom(&ctx);
match atom.node() {
AtomNode::Fun(name, args) => {
assert_eq!(name.as_str(), "T");
assert_eq!(args.len(), 2);
}
_ => panic!("expected Fun node"),
}
}
}