use std::collections::{BTreeMap, BTreeSet};
use egg::{Id, Language};
use laddu_expr::{BinaryOp, UnaryOp, ValueKind};
use num::complex::Complex64;
use super::{Cas, Head, Term, same_shape};
pub(super) struct CandidateBuilder<'a>(pub(super) &'a mut Cas);
impl CandidateBuilder<'_> {
pub(super) fn constant(&mut self, value: Complex64) -> Id {
let head = if value.im == 0.0 && value.im.is_sign_positive() {
Head::Real(value.re.to_bits())
} else {
Head::Complex(value.re.to_bits(), value.im.to_bits())
};
self.0.graph.add(Term::new(head, []))
}
pub(super) fn unary(&mut self, op: UnaryOp, input: Id) -> Option<Id> {
let input = self.0.graph.find(input);
scalar(self.0.graph[input].data.kind)
.then(|| self.0.graph.add(Term::new(Head::Unary(op), [input])))
}
pub(super) fn binary(&mut self, op: BinaryOp, lhs: Id, rhs: Id) -> Option<Id> {
let lhs = self.0.graph.find(lhs);
let rhs = self.0.graph.find(rhs);
(scalar(self.0.graph[lhs].data.kind) && scalar(self.0.graph[rhs].data.kind))
.then(|| self.0.graph.add(Term::new(Head::Binary(op), [lhs, rhs])))
}
pub(super) fn component(&mut self, vector: Id, index: usize) -> Option<Id> {
let vector = self.0.graph.find(vector);
matches!(self.0.graph[vector].data.kind, ValueKind::Vector { len } if index < len).then(
|| {
self.0
.graph
.add(Term::new(Head::Component(index), [vector]))
},
)
}
pub(super) fn matrix_element(&mut self, matrix: Id, row: usize, col: usize) -> Option<Id> {
let matrix = self.0.graph.find(matrix);
matches!(self.0.graph[matrix].data.kind,
ValueKind::Matrix { rows, cols } if row < rows && col < cols)
.then(|| {
self.0
.graph
.add(Term::new(Head::MatrixElement(row, col), [matrix]))
})
}
pub(super) fn vector(&mut self, elements: &[Id]) -> Option<Id> {
let elements = elements
.iter()
.map(|id| self.0.graph.find(*id))
.collect::<Vec<_>>();
elements
.iter()
.all(|id| scalar(self.0.graph[*id].data.kind))
.then(|| self.0.graph.add(Term::new(Head::Vector, elements)))
}
pub(super) fn matrix(&mut self, rows: usize, cols: usize, elements: &[Id]) -> Option<Id> {
if elements.len() != rows.checked_mul(cols)? {
return None;
}
let elements = elements
.iter()
.map(|id| self.0.graph.find(*id))
.collect::<Vec<_>>();
elements
.iter()
.all(|id| scalar(self.0.graph[*id].data.kind))
.then(|| {
self.0
.graph
.add(Term::new(Head::Matrix { rows, cols }, elements))
})
}
pub(super) fn dot(&mut self, lhs: Id, rhs: Id) -> Option<Id> {
let lhs = self.0.graph.find(lhs);
let rhs = self.0.graph.find(rhs);
matches!(
(self.0.graph[lhs].data.kind, self.0.graph[rhs].data.kind),
(ValueKind::Vector { len: a }, ValueKind::Vector { len: b }) if a == b
)
.then(|| self.0.graph.add(Term::new(Head::Dot, [lhs, rhs])))
}
pub(super) fn matvec(&mut self, matrix: Id, vector: Id) -> Option<Id> {
let matrix = self.0.graph.find(matrix);
let vector = self.0.graph.find(vector);
matches!(
(self.0.graph[matrix].data.kind, self.0.graph[vector].data.kind),
(ValueKind::Matrix { cols, .. }, ValueKind::Vector { len }) if cols == len
)
.then(|| self.0.graph.add(Term::new(Head::MatVec, [matrix, vector])))
}
pub(super) fn matmul(&mut self, lhs: Id, rhs: Id) -> Option<Id> {
let lhs = self.0.graph.find(lhs);
let rhs = self.0.graph.find(rhs);
matches!(
(self.0.graph[lhs].data.kind, self.0.graph[rhs].data.kind),
(ValueKind::Matrix { cols, .. }, ValueKind::Matrix { rows, .. }) if cols == rows
)
.then(|| self.0.graph.add(Term::new(Head::MatMul, [lhs, rhs])))
}
pub(super) fn scale(&mut self, value: Id, factor: Id) -> Option<Id> {
let value = self.0.graph.find(value);
let factor = self.0.graph.find(factor);
if !scalar(self.0.graph[factor].data.kind) {
return None;
}
match self.0.graph[value].data.kind {
kind if scalar(kind) => self.product(&[factor, value]),
ValueKind::Vector { len } if len <= 16 => {
let elements = (0..len)
.map(|index| {
let element = self.component(value, index)?;
self.product(&[factor, element])
})
.collect::<Option<Vec<_>>>()?;
self.vector(&elements)
}
ValueKind::Matrix { rows, cols } if rows.checked_mul(cols)? <= 16 => {
let elements = (0..rows)
.flat_map(|row| (0..cols).map(move |col| (row, col)))
.map(|(row, col)| {
let element = self.matrix_element(value, row, col)?;
self.product(&[factor, element])
})
.collect::<Option<Vec<_>>>()?;
self.matrix(rows, cols, &elements)
}
_ => None,
}
}
pub(super) fn sum(&mut self, terms: &[Id]) -> Option<Id> {
self.associative(Head::Sum, terms)
}
pub(super) fn product(&mut self, factors: &[Id]) -> Option<Id> {
self.associative(Head::Product, factors)
}
fn associative(&mut self, head: Head, children: &[Id]) -> Option<Id> {
let mut canonical = children
.iter()
.map(|id| self.0.graph.find(*id))
.collect::<Vec<_>>();
if !canonical
.iter()
.all(|id| scalar(self.0.graph[*id].data.kind))
{
return None;
}
match canonical.as_slice() {
[] => {
return Some(self.constant(if head == Head::Sum {
Complex64::ZERO
} else {
Complex64::ONE
}));
}
[only] => return Some(*only),
_ => (),
}
canonical.sort_unstable();
Some(self.0.graph.add(Term::new(head, canonical)))
}
pub(super) fn emit_equivalent(&mut self, original: Id, candidate: Id) -> bool {
let original = self.0.graph.find(original);
let candidate = self.0.graph.find(candidate);
if !same_shape(
self.0.graph[original].data.kind,
self.0.graph[candidate].data.kind,
) {
return false;
}
self.0.graph.union(original, candidate)
}
}
pub(super) fn scalar(kind: ValueKind) -> bool {
matches!(kind, ValueKind::Real | ValueKind::Complex)
}
pub(super) fn constant(cas: &Cas, id: Id) -> Option<Complex64> {
let id = cas.graph.find(id);
cas.graph[id].nodes.iter().find_map(|node| match node.head {
Head::Real(bits) => Some(Complex64::from(f64::from_bits(bits))),
Head::Complex(re, im) => Some(Complex64::new(f64::from_bits(re), f64::from_bits(im))),
_ => None,
})
}
pub(super) struct ProductView {
pub(super) coefficient: Complex64,
pub(super) real_coefficient: bool,
pub(super) factors: BTreeMap<Id, i32>,
}
impl ProductView {
pub(super) fn of(cas: &Cas, id: Id) -> Self {
let id = cas.graph.find(id);
let mut view = Self {
coefficient: Complex64::ONE,
real_coefficient: true,
factors: BTreeMap::new(),
};
if let Some(product) = cas.graph[id]
.nodes
.iter()
.find(|node| node.head == Head::Product)
{
for &factor in product.children() {
view.push(cas, factor, &mut BTreeSet::new());
}
} else {
view.push(cas, id, &mut BTreeSet::new());
}
view
}
fn push(&mut self, cas: &Cas, id: Id, seen: &mut BTreeSet<Id>) {
let id = cas.graph.find(id);
if !seen.insert(id) {
*self.factors.entry(id).or_default() += 1;
return;
}
if let Some(value) = constant(cas, id) {
if self.real_coefficient
&& cas.graph[id]
.nodes
.iter()
.any(|node| matches!(node.head, Head::Real(_)))
{
self.coefficient = Complex64::from(self.coefficient.re * value.re);
} else {
self.real_coefficient = false;
self.coefficient *= value;
}
return;
}
if !cas.graph[id].nodes.iter().any(|node| {
matches!(
node.head,
Head::Source(_) | Head::Sum | Head::Binary(BinaryOp::Sub | BinaryOp::Add)
)
}) && let Some(nested) = cas.graph[id].nodes.iter().find(|node| {
node.head == Head::Product
&& node
.children()
.iter()
.all(|child| cas.graph.find(*child) != id)
}) {
for &factor in nested.children() {
self.push(cas, factor, seen);
}
return;
}
if let Some(power) = cas.graph[id].nodes.iter().find_map(|node| match node.head {
Head::Unary(UnaryOp::PowI(power)) => Some((node.children[0], power)),
_ => None,
}) {
*self.factors.entry(cas.graph.find(power.0)).or_default() += power.1;
} else if !cas.graph[id]
.nodes
.iter()
.any(|node| matches!(node.head, Head::Source(_)))
&& let Some(input) = cas.graph[id].nodes.iter().find_map(|node| match node.head {
Head::Unary(UnaryOp::Neg) => Some(node.children[0]),
_ => None,
})
{
self.coefficient = if self.real_coefficient {
Complex64::from(-self.coefficient.re)
} else {
-self.coefficient
};
self.push(cas, input, seen);
} else {
*self.factors.entry(id).or_default() += 1;
}
}
pub(super) fn emit(&self, builder: &mut CandidateBuilder<'_>) -> Option<Id> {
if self.coefficient == Complex64::ONE
&& self.factors.len() > 1
&& let Some(&common_power) = self.factors.values().next()
&& common_power > 1
&& self.factors.values().all(|power| *power == common_power)
{
let bases = self.factors.keys().copied().collect::<Vec<_>>();
let base = builder.product(&bases)?;
return builder.unary(UnaryOp::PowI(common_power), base);
}
let mut factors = Vec::new();
if self.coefficient != Complex64::ONE || self.factors.is_empty() {
factors.push(builder.constant(self.coefficient));
}
for (&id, &power) in &self.factors {
if power == 0 {
continue;
}
factors.push(if power == 1 {
id
} else {
builder.unary(UnaryOp::PowI(power), id)?
});
}
builder.product(&factors)
}
}