use super::LoweredOp;
use std::collections::HashMap;
use std::sync::Arc;
#[inline(always)]
fn pop_or_nan(stack: &mut Vec<f64>) -> f64 {
debug_assert!(!stack.is_empty(), "OxiOp stack underflow — malformed IR");
stack.pop().unwrap_or(f64::NAN)
}
#[derive(Clone, Debug, PartialEq)]
pub enum OxiOp {
Const(f64),
Var(usize),
Add,
Sub,
Mul,
Div,
Neg,
Exp,
Ln,
Sin,
Cos,
Pow,
Tan,
Sinh,
Cosh,
Tanh,
Arcsin,
Arccos,
Arctan,
Arcsinh,
Arccosh,
Arctanh,
Erf,
LGamma,
Digamma,
Trigamma,
Ei,
Si,
Ci,
Store(usize),
Load(usize),
}
impl LoweredOp {
pub fn to_oxiblas_ops(&self) -> Vec<OxiOp> {
let mut ops = Vec::new();
self.collect_ops(&mut ops);
ops
}
pub(super) fn collect_ops(&self, ops: &mut Vec<OxiOp>) {
match self {
Self::Const(c) => ops.push(OxiOp::Const(*c)),
Self::NamedConst(nc) => ops.push(OxiOp::Const(nc.value())),
Self::Var(i) => ops.push(OxiOp::Var(*i)),
Self::Add(a, b) => {
a.collect_ops(ops);
b.collect_ops(ops);
ops.push(OxiOp::Add);
}
Self::Sub(a, b) => {
a.collect_ops(ops);
b.collect_ops(ops);
ops.push(OxiOp::Sub);
}
Self::Mul(a, b) => {
a.collect_ops(ops);
b.collect_ops(ops);
ops.push(OxiOp::Mul);
}
Self::Div(a, b) => {
a.collect_ops(ops);
b.collect_ops(ops);
ops.push(OxiOp::Div);
}
Self::Exp(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Exp);
}
Self::Ln(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Ln);
}
Self::Sin(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Sin);
}
Self::Cos(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Cos);
}
Self::Pow(a, b) => {
a.collect_ops(ops);
b.collect_ops(ops);
ops.push(OxiOp::Pow);
}
Self::Neg(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Neg);
}
Self::Tan(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Tan);
}
Self::Sinh(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Sinh);
}
Self::Cosh(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Cosh);
}
Self::Tanh(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Tanh);
}
Self::Arcsin(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Arcsin);
}
Self::Arccos(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Arccos);
}
Self::Arctan(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Arctan);
}
Self::Arcsinh(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Arcsinh);
}
Self::Arccosh(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Arccosh);
}
Self::Arctanh(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Arctanh);
}
Self::Erf(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Erf);
}
Self::LGamma(a) => {
a.collect_ops(ops);
ops.push(OxiOp::LGamma);
}
Self::Digamma(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Digamma);
}
Self::Trigamma(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Trigamma);
}
Self::Ei(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Ei);
}
Self::Si(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Si);
}
Self::Ci(a) => {
a.collect_ops(ops);
ops.push(OxiOp::Ci);
}
}
}
pub fn eval_ops(ops: &[OxiOp], vars: &[f64]) -> f64 {
let mut stack: Vec<f64> = Vec::with_capacity(ops.len());
for op in ops {
match op {
OxiOp::Const(c) => stack.push(*c),
OxiOp::Var(i) => {
stack.push(vars.get(*i).copied().unwrap_or(f64::NAN));
}
OxiOp::Add => {
let b = pop_or_nan(&mut stack);
let a = pop_or_nan(&mut stack);
stack.push(a + b);
}
OxiOp::Sub => {
let b = pop_or_nan(&mut stack);
let a = pop_or_nan(&mut stack);
stack.push(a - b);
}
OxiOp::Mul => {
let b = pop_or_nan(&mut stack);
let a = pop_or_nan(&mut stack);
stack.push(a * b);
}
OxiOp::Div => {
let b = pop_or_nan(&mut stack);
let a = pop_or_nan(&mut stack);
stack.push(a / b);
}
OxiOp::Neg => {
let a = pop_or_nan(&mut stack);
stack.push(-a);
}
OxiOp::Exp => {
let a = pop_or_nan(&mut stack);
stack.push(a.exp());
}
OxiOp::Ln => {
let a = pop_or_nan(&mut stack);
stack.push(a.ln());
}
OxiOp::Sin => {
let a = pop_or_nan(&mut stack);
stack.push(a.sin());
}
OxiOp::Cos => {
let a = pop_or_nan(&mut stack);
stack.push(a.cos());
}
OxiOp::Pow => {
let b = pop_or_nan(&mut stack);
let a = pop_or_nan(&mut stack);
stack.push(a.powf(b));
}
OxiOp::Tan => {
let a = pop_or_nan(&mut stack);
stack.push(a.tan());
}
OxiOp::Sinh => {
let a = pop_or_nan(&mut stack);
stack.push(a.sinh());
}
OxiOp::Cosh => {
let a = pop_or_nan(&mut stack);
stack.push(a.cosh());
}
OxiOp::Tanh => {
let a = pop_or_nan(&mut stack);
stack.push(a.tanh());
}
OxiOp::Arcsin => {
let a = pop_or_nan(&mut stack);
stack.push(a.asin());
}
OxiOp::Arccos => {
let a = pop_or_nan(&mut stack);
stack.push(a.acos());
}
OxiOp::Arctan => {
let a = pop_or_nan(&mut stack);
stack.push(a.atan());
}
OxiOp::Arcsinh => {
let a = pop_or_nan(&mut stack);
stack.push(a.asinh());
}
OxiOp::Arccosh => {
let a = pop_or_nan(&mut stack);
stack.push(a.acosh());
}
OxiOp::Arctanh => {
let a = pop_or_nan(&mut stack);
stack.push(a.atanh());
}
OxiOp::Erf => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::erf(a));
}
OxiOp::LGamma => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::lgamma(a));
}
OxiOp::Digamma => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::digamma(a));
}
OxiOp::Trigamma => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::trigamma(a));
}
OxiOp::Ei => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::ei(a));
}
OxiOp::Si => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::si(a));
}
OxiOp::Ci => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::ci(a));
}
OxiOp::Store(_) => {
debug_assert!(
false,
"OxiOp::Store in eval_ops — use eval_ops_shared instead"
);
}
OxiOp::Load(_) => {
debug_assert!(
false,
"OxiOp::Load in eval_ops — use eval_ops_shared instead"
);
stack.push(f64::NAN);
}
}
}
pop_or_nan(&mut stack)
}
pub fn eval_batch(&self, data: &[Vec<f64>]) -> Vec<f64> {
let ops = self.to_oxiblas_ops();
#[cfg(feature = "simd")]
{
crate::simd_eval::eval_batch_simd(&ops, data)
}
#[cfg(not(feature = "simd"))]
{
Self::eval_batch_scalar_from_ops(&ops, data)
}
}
pub fn eval_batch_scalar_from_ops(ops: &[OxiOp], data: &[Vec<f64>]) -> Vec<f64> {
data.iter().map(|row| Self::eval_ops(ops, row)).collect()
}
pub fn eval_batch_scalar(&self, data: &[Vec<f64>]) -> Vec<f64> {
let ops = self.to_oxiblas_ops();
Self::eval_batch_scalar_from_ops(&ops, data)
}
pub fn to_oxiblas_ops_shared(&self) -> (Vec<OxiOp>, usize) {
let root = Arc::new(self.clone());
let mut refcount: HashMap<*const LoweredOp, u32> = HashMap::new();
census(&root, &mut refcount);
let mut slot_for: HashMap<*const LoweredOp, usize> = HashMap::new();
let mut next_slot = 0usize;
for (&ptr, &count) in &refcount {
if count >= 2 {
slot_for.insert(ptr, next_slot);
next_slot += 1;
}
}
let n_slots = next_slot;
let mut ops = Vec::new();
let mut emitted: HashMap<*const LoweredOp, usize> = HashMap::new();
emit_shared(&root, &slot_for, &mut emitted, &mut ops);
(ops, n_slots)
}
pub fn eval_ops_shared(ops: &[OxiOp], vars: &[f64], n_slots: usize) -> f64 {
let mut stack: Vec<f64> = Vec::with_capacity(ops.len());
let mut slots: Vec<f64> = vec![f64::NAN; n_slots];
for op in ops {
match op {
OxiOp::Store(k) => {
debug_assert!(
!stack.is_empty(),
"OxiOp::Store({k}) stack underflow — malformed IR"
);
let v = stack.last().copied().unwrap_or(f64::NAN);
if let Some(slot) = slots.get_mut(*k) {
*slot = v;
} else {
debug_assert!(
false,
"OxiOp::Store({k}) slot index out of range (n_slots={n_slots})"
);
}
}
OxiOp::Load(k) => {
let v = slots.get(*k).copied().unwrap_or(f64::NAN);
debug_assert!(
*k < n_slots,
"OxiOp::Load({k}) slot index out of range (n_slots={n_slots})"
);
stack.push(v);
}
OxiOp::Const(c) => stack.push(*c),
OxiOp::Var(i) => {
stack.push(vars.get(*i).copied().unwrap_or(f64::NAN));
}
OxiOp::Add => {
let b = pop_or_nan(&mut stack);
let a = pop_or_nan(&mut stack);
stack.push(a + b);
}
OxiOp::Sub => {
let b = pop_or_nan(&mut stack);
let a = pop_or_nan(&mut stack);
stack.push(a - b);
}
OxiOp::Mul => {
let b = pop_or_nan(&mut stack);
let a = pop_or_nan(&mut stack);
stack.push(a * b);
}
OxiOp::Div => {
let b = pop_or_nan(&mut stack);
let a = pop_or_nan(&mut stack);
stack.push(a / b);
}
OxiOp::Neg => {
let a = pop_or_nan(&mut stack);
stack.push(-a);
}
OxiOp::Exp => {
let a = pop_or_nan(&mut stack);
stack.push(a.exp());
}
OxiOp::Ln => {
let a = pop_or_nan(&mut stack);
stack.push(a.ln());
}
OxiOp::Sin => {
let a = pop_or_nan(&mut stack);
stack.push(a.sin());
}
OxiOp::Cos => {
let a = pop_or_nan(&mut stack);
stack.push(a.cos());
}
OxiOp::Pow => {
let b = pop_or_nan(&mut stack);
let a = pop_or_nan(&mut stack);
stack.push(a.powf(b));
}
OxiOp::Tan => {
let a = pop_or_nan(&mut stack);
stack.push(a.tan());
}
OxiOp::Sinh => {
let a = pop_or_nan(&mut stack);
stack.push(a.sinh());
}
OxiOp::Cosh => {
let a = pop_or_nan(&mut stack);
stack.push(a.cosh());
}
OxiOp::Tanh => {
let a = pop_or_nan(&mut stack);
stack.push(a.tanh());
}
OxiOp::Arcsin => {
let a = pop_or_nan(&mut stack);
stack.push(a.asin());
}
OxiOp::Arccos => {
let a = pop_or_nan(&mut stack);
stack.push(a.acos());
}
OxiOp::Arctan => {
let a = pop_or_nan(&mut stack);
stack.push(a.atan());
}
OxiOp::Arcsinh => {
let a = pop_or_nan(&mut stack);
stack.push(a.asinh());
}
OxiOp::Arccosh => {
let a = pop_or_nan(&mut stack);
stack.push(a.acosh());
}
OxiOp::Arctanh => {
let a = pop_or_nan(&mut stack);
stack.push(a.atanh());
}
OxiOp::Erf => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::erf(a));
}
OxiOp::LGamma => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::lgamma(a));
}
OxiOp::Digamma => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::digamma(a));
}
OxiOp::Trigamma => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::trigamma(a));
}
OxiOp::Ei => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::ei(a));
}
OxiOp::Si => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::si(a));
}
OxiOp::Ci => {
let a = pop_or_nan(&mut stack);
stack.push(crate::special::ci(a));
}
}
}
pop_or_nan(&mut stack)
}
}
fn census(node: &Arc<LoweredOp>, counts: &mut HashMap<*const LoweredOp, u32>) {
let ptr = Arc::as_ptr(node);
let c = counts.entry(ptr).or_insert(0);
*c += 1;
if *c > 1 {
return;
}
match node.as_ref() {
LoweredOp::Const(_) | LoweredOp::Var(_) | LoweredOp::NamedConst(_) => {}
LoweredOp::Neg(a)
| LoweredOp::Exp(a)
| LoweredOp::Ln(a)
| LoweredOp::Sin(a)
| LoweredOp::Cos(a)
| LoweredOp::Tan(a)
| LoweredOp::Sinh(a)
| LoweredOp::Cosh(a)
| LoweredOp::Tanh(a)
| LoweredOp::Arcsin(a)
| LoweredOp::Arccos(a)
| LoweredOp::Arctan(a)
| LoweredOp::Arcsinh(a)
| LoweredOp::Arccosh(a)
| LoweredOp::Arctanh(a)
| LoweredOp::Erf(a)
| LoweredOp::LGamma(a)
| LoweredOp::Digamma(a)
| LoweredOp::Trigamma(a)
| LoweredOp::Ei(a)
| LoweredOp::Si(a)
| LoweredOp::Ci(a) => {
census(a, counts);
}
LoweredOp::Add(a, b)
| LoweredOp::Sub(a, b)
| LoweredOp::Mul(a, b)
| LoweredOp::Div(a, b)
| LoweredOp::Pow(a, b) => {
census(a, counts);
census(b, counts);
}
}
}
fn emit_shared(
node: &Arc<LoweredOp>,
slot_for: &HashMap<*const LoweredOp, usize>,
emitted: &mut HashMap<*const LoweredOp, usize>,
ops: &mut Vec<OxiOp>,
) {
let ptr = Arc::as_ptr(node);
if let Some(&slot) = emitted.get(&ptr) {
ops.push(OxiOp::Load(slot));
return;
}
match node.as_ref() {
LoweredOp::Const(_) | LoweredOp::Var(_) | LoweredOp::NamedConst(_) => {}
LoweredOp::Neg(a)
| LoweredOp::Exp(a)
| LoweredOp::Ln(a)
| LoweredOp::Sin(a)
| LoweredOp::Cos(a)
| LoweredOp::Tan(a)
| LoweredOp::Sinh(a)
| LoweredOp::Cosh(a)
| LoweredOp::Tanh(a)
| LoweredOp::Arcsin(a)
| LoweredOp::Arccos(a)
| LoweredOp::Arctan(a)
| LoweredOp::Arcsinh(a)
| LoweredOp::Arccosh(a)
| LoweredOp::Arctanh(a)
| LoweredOp::Erf(a)
| LoweredOp::LGamma(a)
| LoweredOp::Digamma(a)
| LoweredOp::Trigamma(a)
| LoweredOp::Ei(a)
| LoweredOp::Si(a)
| LoweredOp::Ci(a) => {
emit_shared(a, slot_for, emitted, ops);
}
LoweredOp::Add(a, b)
| LoweredOp::Sub(a, b)
| LoweredOp::Mul(a, b)
| LoweredOp::Div(a, b)
| LoweredOp::Pow(a, b) => {
emit_shared(a, slot_for, emitted, ops);
emit_shared(b, slot_for, emitted, ops);
}
}
match node.as_ref() {
LoweredOp::Const(c) => ops.push(OxiOp::Const(*c)),
LoweredOp::NamedConst(nc) => ops.push(OxiOp::Const(nc.value())),
LoweredOp::Var(i) => ops.push(OxiOp::Var(*i)),
LoweredOp::Add(_, _) => ops.push(OxiOp::Add),
LoweredOp::Sub(_, _) => ops.push(OxiOp::Sub),
LoweredOp::Mul(_, _) => ops.push(OxiOp::Mul),
LoweredOp::Div(_, _) => ops.push(OxiOp::Div),
LoweredOp::Exp(_) => ops.push(OxiOp::Exp),
LoweredOp::Ln(_) => ops.push(OxiOp::Ln),
LoweredOp::Sin(_) => ops.push(OxiOp::Sin),
LoweredOp::Cos(_) => ops.push(OxiOp::Cos),
LoweredOp::Pow(_, _) => ops.push(OxiOp::Pow),
LoweredOp::Neg(_) => ops.push(OxiOp::Neg),
LoweredOp::Tan(_) => ops.push(OxiOp::Tan),
LoweredOp::Sinh(_) => ops.push(OxiOp::Sinh),
LoweredOp::Cosh(_) => ops.push(OxiOp::Cosh),
LoweredOp::Tanh(_) => ops.push(OxiOp::Tanh),
LoweredOp::Arcsin(_) => ops.push(OxiOp::Arcsin),
LoweredOp::Arccos(_) => ops.push(OxiOp::Arccos),
LoweredOp::Arctan(_) => ops.push(OxiOp::Arctan),
LoweredOp::Arcsinh(_) => ops.push(OxiOp::Arcsinh),
LoweredOp::Arccosh(_) => ops.push(OxiOp::Arccosh),
LoweredOp::Arctanh(_) => ops.push(OxiOp::Arctanh),
LoweredOp::Erf(_) => ops.push(OxiOp::Erf),
LoweredOp::LGamma(_) => ops.push(OxiOp::LGamma),
LoweredOp::Digamma(_) => ops.push(OxiOp::Digamma),
LoweredOp::Trigamma(_) => ops.push(OxiOp::Trigamma),
LoweredOp::Ei(_) => ops.push(OxiOp::Ei),
LoweredOp::Si(_) => ops.push(OxiOp::Si),
LoweredOp::Ci(_) => ops.push(OxiOp::Ci),
}
if let Some(&slot) = slot_for.get(&ptr) {
ops.push(OxiOp::Store(slot));
emitted.insert(ptr, slot);
}
}