pub mod display;
pub mod oxiblas;
pub mod pattern;
pub use crate::lower_interval::IntervalLO;
pub use oxiblas::OxiOp;
use crate::error::EmlError;
use crate::eval::EvalCtx;
use crate::named_const::NamedConst;
use crate::tree::EmlTree;
use std::sync::Arc;
#[derive(Clone, Debug, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(rename_all = "snake_case"))]
pub enum LoweredOp {
Const(f64),
Var(usize),
Add(Arc<LoweredOp>, Arc<LoweredOp>),
Sub(Arc<LoweredOp>, Arc<LoweredOp>),
Mul(Arc<LoweredOp>, Arc<LoweredOp>),
Div(Arc<LoweredOp>, Arc<LoweredOp>),
Exp(Arc<LoweredOp>),
Ln(Arc<LoweredOp>),
Sin(Arc<LoweredOp>),
Cos(Arc<LoweredOp>),
Pow(Arc<LoweredOp>, Arc<LoweredOp>),
Neg(Arc<LoweredOp>),
Tan(Arc<LoweredOp>),
Sinh(Arc<LoweredOp>),
Cosh(Arc<LoweredOp>),
Tanh(Arc<LoweredOp>),
Arcsin(Arc<LoweredOp>),
Arccos(Arc<LoweredOp>),
Arctan(Arc<LoweredOp>),
Arcsinh(Arc<LoweredOp>),
Arccosh(Arc<LoweredOp>),
Arctanh(Arc<LoweredOp>),
NamedConst(NamedConst),
Erf(Arc<LoweredOp>),
LGamma(Arc<LoweredOp>),
Digamma(Arc<LoweredOp>),
Trigamma(Arc<LoweredOp>),
Ei(Arc<LoweredOp>),
Si(Arc<LoweredOp>),
Ci(Arc<LoweredOp>),
}
impl EmlTree {
pub fn lower(&self) -> LoweredOp {
pattern::lower_node(&self.root)
}
pub fn eval_real_lowered(&self, ctx: &EvalCtx) -> Result<f64, EmlError> {
let cse_root = self.lower().simplify().cse();
let (ops, n_slots) = cse_root.to_oxiblas_ops_shared();
let result = LoweredOp::eval_ops_shared(&ops, ctx.as_slice(), n_slots);
if result.is_nan() {
return Err(EmlError::NanEncountered);
}
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_lower_one() {
let t = EmlTree::one();
let lowered = t.lower();
assert_eq!(lowered, LoweredOp::Const(1.0));
}
#[test]
fn test_lower_var() {
let t = EmlTree::var(0);
let lowered = t.lower();
assert_eq!(lowered, LoweredOp::Var(0));
}
#[test]
fn test_lower_exp() {
let x = EmlTree::var(0);
let one = EmlTree::one();
let exp_x = EmlTree::eml(&x, &one);
let lowered = exp_x.lower();
assert_eq!(lowered, LoweredOp::Exp(Arc::new(LoweredOp::Var(0))));
}
#[test]
fn test_lower_e_minus_x() {
let x = EmlTree::var(0);
let one = EmlTree::one();
let exp_x = EmlTree::eml(&x, &one);
let e_minus_x = EmlTree::eml(&one, &exp_x);
let lowered = e_minus_x.lower();
assert_eq!(
lowered,
LoweredOp::Sub(
Arc::new(LoweredOp::Const(std::f64::consts::E)),
Arc::new(LoweredOp::Var(0)),
)
);
}
#[test]
fn test_lower_ln() {
let x = EmlTree::var(0);
let one = EmlTree::one();
let inner = EmlTree::eml(&one, &x); let middle = EmlTree::eml(&inner, &one); let ln_x = EmlTree::eml(&one, &middle); let lowered = ln_x.lower();
assert_eq!(lowered, LoweredOp::Ln(Arc::new(LoweredOp::Var(0))));
}
#[test]
fn test_lowered_eval() {
let op = LoweredOp::Add(Arc::new(LoweredOp::Var(0)), Arc::new(LoweredOp::Const(3.0)));
assert!((op.eval(&[2.0]) - 5.0).abs() < 1e-15);
}
#[test]
fn test_pretty_print() {
let op = LoweredOp::Mul(Arc::new(LoweredOp::Var(0)), Arc::new(LoweredOp::Var(1)));
assert_eq!(op.to_pretty(), "(x0 * x1)");
}
#[test]
fn test_simplify_exp_ln() {
let op = LoweredOp::Exp(Arc::new(LoweredOp::Ln(Arc::new(LoweredOp::Var(0)))));
let simplified = op.simplify();
assert_eq!(simplified, LoweredOp::Var(0));
}
#[test]
fn test_simplify_constants() {
let op = LoweredOp::Add(
Arc::new(LoweredOp::Const(2.0)),
Arc::new(LoweredOp::Const(3.0)),
);
let simplified = op.simplify();
assert_eq!(simplified, LoweredOp::Const(5.0));
}
#[test]
fn test_to_oxiblas_ops_roundtrip() {
use crate::Canonical;
let x = crate::tree::EmlTree::var(0);
let exp_x = Canonical::exp(&x);
let lowered = exp_x.lower();
let ops = lowered.to_oxiblas_ops();
let result = LoweredOp::eval_ops(&ops, &[1.5_f64]);
assert!(
(result - 1.5_f64.exp()).abs() < 1e-12,
"exp roundtrip failed: {result}"
);
let ln_x = Canonical::ln(&x);
let lowered_ln = ln_x.lower();
let ops_ln = lowered_ln.to_oxiblas_ops();
let result_ln = LoweredOp::eval_ops(&ops_ln, &[2.0_f64]);
assert!(
(result_ln - 2.0_f64.ln()).abs() < 1e-12,
"ln roundtrip failed: {result_ln}"
);
let lowered_sin = LoweredOp::Sin(Arc::new(LoweredOp::Var(0)));
let ops_sin = lowered_sin.to_oxiblas_ops();
let result_sin = LoweredOp::eval_ops(&ops_sin, &[std::f64::consts::PI / 6.0]);
assert!(
(result_sin - 0.5_f64).abs() < 1e-9,
"sin roundtrip failed: {result_sin}"
);
}
#[test]
fn test_eval_batch_scalar_matches_eval() {
use crate::Canonical;
let x = crate::tree::EmlTree::var(0);
let exp_x = Canonical::exp(&x);
let lowered = exp_x.lower();
let data: Vec<Vec<f64>> = (0..100).map(|i| vec![i as f64 * 0.05]).collect();
let batch_results = lowered.eval_batch_scalar(&data);
assert_eq!(batch_results.len(), 100);
for (row, result) in data.iter().zip(batch_results.iter()) {
let expected = lowered.eval(row);
assert!(
(result - expected).abs() < 1e-12,
"mismatch at x={}: got {result}, expected {expected}",
row[0]
);
}
}
#[test]
fn test_structural_hash_differs() {
use crate::Canonical;
use std::collections::hash_map::DefaultHasher;
use std::hash::Hasher;
let x = crate::tree::EmlTree::var(0);
let exp_x = Canonical::exp(&x).lower().simplify();
let ln_x = Canonical::ln(&x).lower().simplify();
let mut h1 = DefaultHasher::new();
exp_x.structural_hash(&mut h1);
let mut h2 = DefaultHasher::new();
ln_x.structural_hash(&mut h2);
assert_ne!(
h1.finish(),
h2.finish(),
"exp and ln should have different structural hashes"
);
}
#[test]
fn test_structural_hash_same_for_equiv() {
use crate::Canonical;
use std::collections::hash_map::DefaultHasher;
use std::hash::Hasher;
let x = crate::tree::EmlTree::var(0);
let exp_x1 = Canonical::exp(&x).lower().simplify();
let exp_x2 = Canonical::exp(&x).lower().simplify();
let mut h1 = DefaultHasher::new();
exp_x1.structural_hash(&mut h1);
let mut h2 = DefaultHasher::new();
exp_x2.structural_hash(&mut h2);
assert_eq!(
h1.finish(),
h2.finish(),
"identical trees should have the same structural hash"
);
}
#[test]
fn latex_var() {
assert_eq!(LoweredOp::Var(0).to_latex(), "x_{0}");
assert_eq!(LoweredOp::Var(3).to_latex(), "x_{3}");
}
#[test]
fn latex_const_pi() {
assert_eq!(LoweredOp::Const(std::f64::consts::PI).to_latex(), r"\pi");
}
#[test]
fn latex_const_e() {
assert_eq!(LoweredOp::Const(std::f64::consts::E).to_latex(), "e");
}
#[test]
fn latex_const_integer() {
assert_eq!(LoweredOp::Const(2.0).to_latex(), "2");
assert_eq!(LoweredOp::Const(-1.0).to_latex(), "-1");
}
#[test]
fn latex_div() {
let op = LoweredOp::Div(Arc::new(LoweredOp::Const(1.0)), Arc::new(LoweredOp::Var(0)));
assert_eq!(op.to_latex(), r"\frac{1}{x_{0}}");
}
#[test]
fn latex_exp() {
let op = LoweredOp::Exp(Arc::new(LoweredOp::Var(0)));
assert_eq!(op.to_latex(), r"e^{x_{0}}");
}
#[test]
fn latex_ln() {
let op = LoweredOp::Ln(Arc::new(LoweredOp::Var(0)));
assert_eq!(op.to_latex(), r"\ln\left(x_{0}\right)");
}
#[test]
fn latex_sin_cos() {
let op = LoweredOp::Sin(Arc::new(LoweredOp::Var(0)));
assert_eq!(op.to_latex(), r"\sin\left(x_{0}\right)");
let op2 = LoweredOp::Cos(Arc::new(LoweredOp::Var(0)));
assert_eq!(op2.to_latex(), r"\cos\left(x_{0}\right)");
}
#[test]
fn latex_pow() {
let op = LoweredOp::Pow(Arc::new(LoweredOp::Var(0)), Arc::new(LoweredOp::Const(2.0)));
assert_eq!(op.to_latex(), "x_{0}^{2}");
}
#[test]
fn latex_neg() {
let op = LoweredOp::Neg(Arc::new(LoweredOp::Var(0)));
assert_eq!(op.to_latex(), "-x_{0}");
}
#[test]
fn latex_mul() {
let op = LoweredOp::Mul(Arc::new(LoweredOp::Const(2.0)), Arc::new(LoweredOp::Var(0)));
assert_eq!(op.to_latex(), r"2 \cdot x_{0}");
}
#[test]
fn latex_composite() {
let op = LoweredOp::Div(
Arc::new(LoweredOp::Sin(Arc::new(LoweredOp::Var(0)))),
Arc::new(LoweredOp::Cos(Arc::new(LoweredOp::Var(0)))),
);
let latex = op.to_latex();
assert!(latex.contains(r"\frac"));
assert!(latex.contains(r"\sin"));
assert!(latex.contains(r"\cos"));
}
}