use crate::kernel::{ExprData, ExprId, ExprPool};
use crate::poly::{ConversionError, UniPoly};
use std::collections::HashMap;
pub fn horner(expr: ExprId, var: ExprId, pool: &ExprPool) -> Result<ExprId, ConversionError> {
let poly = UniPoly::from_symbolic(expr, var, pool)?;
let coeffs = poly.coefficients_i64(); Ok(build_horner(&coeffs, var, pool))
}
fn build_horner(coeffs: &[i64], var: ExprId, pool: &ExprPool) -> ExprId {
if coeffs.is_empty() {
return pool.integer(0_i32);
}
let n = coeffs.len();
let mut result = pool.integer(coeffs[n - 1]);
for k in (0..n - 1).rev() {
let xr = pool.mul(vec![var, result]);
let ck = pool.integer(coeffs[k]);
result = pool.add(vec![ck, xr]);
}
result
}
pub fn emit_horner_c(
expr: ExprId,
var: ExprId,
var_name: &str,
fn_name: &str,
pool: &ExprPool,
) -> Result<String, ConversionError> {
let poly = UniPoly::from_symbolic(expr, var, pool)?;
let coeffs = poly.coefficients_i64();
let body = build_c_horner(&coeffs, var_name);
Ok(format!(
"double {}(double {}) {{\n return {};\n}}\n",
fn_name, var_name, body
))
}
#[inline]
pub fn eval_horner_f64(coeffs: &[f64], x: f64) -> f64 {
if coeffs.is_empty() {
return 0.0;
}
let mut acc = coeffs[coeffs.len() - 1];
for &c in coeffs[..coeffs.len() - 1].iter().rev() {
acc = c + x * acc;
}
acc
}
pub fn eval_horner_f64_batch(coeffs: &[f64], xs: &[f64], out: &mut [f64]) {
assert_eq!(xs.len(), out.len());
let mut i = 0;
while i + 4 <= xs.len() {
let chunk = wide::f64x4::new([xs[i], xs[i + 1], xs[i + 2], xs[i + 3]]);
let vals = eval_horner_f64x4(coeffs, chunk).to_array();
out[i..i + 4].copy_from_slice(&vals);
i += 4;
}
for (x, o) in xs[i..].iter().zip(out[i..].iter_mut()) {
*o = eval_horner_f64(coeffs, *x);
}
}
#[inline]
fn eval_horner_f64x4(coeffs: &[f64], x: wide::f64x4) -> wide::f64x4 {
if coeffs.is_empty() {
return wide::f64x4::splat(0.0);
}
let mut acc = wide::f64x4::splat(coeffs[coeffs.len() - 1]);
for &c in coeffs[..coeffs.len() - 1].iter().rev() {
acc = wide::f64x4::splat(c) + x * acc;
}
acc
}
fn build_c_horner(coeffs: &[i64], var: &str) -> String {
if coeffs.is_empty() {
return "0.0".to_string();
}
let n = coeffs.len();
let mut result = format!("{}.0", coeffs[n - 1]);
for k in (0..n - 1).rev() {
let ck = format!("{}.0", coeffs[k]);
result = format!("{} + {} * ({})", ck, var, result);
}
result
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum EmitCError {
UnsupportedFunction(String),
UnsupportedNode(String),
MissingVariable(String),
}
impl std::fmt::Display for EmitCError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
EmitCError::UnsupportedFunction(name) => {
write!(f, "function '{name}' has no C math.h equivalent")
}
EmitCError::UnsupportedNode(desc) => {
write!(f, "expression node not supported in C emission: {desc}")
}
EmitCError::MissingVariable(name) => {
write!(
f,
"symbol '{name}' is not listed in the vars/var_names parameter"
)
}
}
}
}
impl std::error::Error for EmitCError {}
fn func_to_c(name: &str, arg_exprs: &[String]) -> Option<String> {
if name == "atan2" && arg_exprs.len() == 2 {
return Some(format!("atan2({}, {})", arg_exprs[0], arg_exprs[1]));
}
if name == "pow" && arg_exprs.len() == 2 {
return Some(format!("pow({}, {})", arg_exprs[0], arg_exprs[1]));
}
if name == "min" && arg_exprs.len() == 2 {
return Some(format!("fmin({}, {})", arg_exprs[0], arg_exprs[1]));
}
if name == "max" && arg_exprs.len() == 2 {
return Some(format!("fmax({}, {})", arg_exprs[0], arg_exprs[1]));
}
if arg_exprs.len() == 1 {
let a = &arg_exprs[0];
let c_name = match name {
"sin" => "sin",
"cos" => "cos",
"tan" => "tan",
"asin" => "asin",
"acos" => "acos",
"atan" => "atan",
"sinh" => "sinh",
"cosh" => "cosh",
"tanh" => "tanh",
"exp" => "exp",
"log" => "log",
"sqrt" => "sqrt",
"abs" => "fabs",
"erf" => "erf",
"erfc" => "erfc",
"tgamma" | "gamma" => "tgamma",
"floor" => "floor",
"ceil" => "ceil",
"round" => "round",
"sign" => return Some(format!("(({a}) > 0.0 ? 1.0 : (({a}) < 0.0 ? -1.0 : 0.0))")),
"log10" => "log10",
"log2" => "log2",
"exp2" => "exp2",
"cbrt" => "cbrt",
_ => return None,
};
return Some(format!("{c_name}({a})"));
}
None
}
fn emit_expr_inner(
expr: ExprId,
var_map: &HashMap<ExprId, &str>,
pool: &ExprPool,
stmts: &mut Vec<String>,
memo: &mut HashMap<ExprId, String>,
counter: &mut usize,
) -> Result<String, EmitCError> {
if let Some(cached) = memo.get(&expr) {
return Ok(cached.clone());
}
if let Some(&name) = var_map.get(&expr) {
return Ok(name.to_string());
}
let result = match pool.get(expr) {
ExprData::Integer(n) => {
format!("{}.0", n.0)
}
ExprData::Rational(r) => {
let v = r.0.numer().to_f64() / r.0.denom().to_f64();
format!("{v:?}")
}
ExprData::Float(f) => {
let v = f.inner.to_f64();
format!("{v:?}")
}
ExprData::Symbol { name, .. } => {
return Err(EmitCError::MissingVariable(name.clone()));
}
ExprData::Add(args) => {
let mut parts = Vec::with_capacity(args.len());
for &a in &args {
parts.push(emit_expr_inner(a, var_map, pool, stmts, memo, counter)?);
}
format!("({})", parts.join(" + "))
}
ExprData::Mul(args) => {
let mut parts = Vec::with_capacity(args.len());
for &a in &args {
parts.push(emit_expr_inner(a, var_map, pool, stmts, memo, counter)?);
}
format!("({})", parts.join(" * "))
}
ExprData::Pow { base, exp } => {
let b = emit_expr_inner(base, var_map, pool, stmts, memo, counter)?;
let e = emit_expr_inner(exp, var_map, pool, stmts, memo, counter)?;
match pool.get(exp) {
ExprData::Integer(ref n) => {
if let Some(k) = n.0.to_i32() {
match k {
0 => "1.0".to_string(),
1 => b,
2 => format!("({b} * {b})"),
3 => format!("({b} * {b} * {b})"),
-1 => format!("(1.0 / {b})"),
_ => format!("pow({b}, {e})"),
}
} else {
format!("pow({b}, {e})")
}
}
_ => format!("pow({b}, {e})"),
}
}
ExprData::Func { name, args } => {
let mut arg_exprs = Vec::with_capacity(args.len());
for &a in &args {
arg_exprs.push(emit_expr_inner(a, var_map, pool, stmts, memo, counter)?);
}
func_to_c(&name, &arg_exprs).ok_or(EmitCError::UnsupportedFunction(name))?
}
ExprData::Piecewise { branches, default } => {
let default_str = emit_expr_inner(default, var_map, pool, stmts, memo, counter)?;
let mut result = default_str;
for (cond, val) in branches.iter().rev() {
let cond_str = emit_predicate_inner(*cond, var_map, pool, stmts, memo, counter)?;
let val_str = emit_expr_inner(*val, var_map, pool, stmts, memo, counter)?;
result = format!("(({cond_str}) ? ({val_str}) : ({result}))");
}
result
}
other => {
return Err(EmitCError::UnsupportedNode(format!("{other:?}")));
}
};
let is_simple = result.len() <= 24
|| result.starts_with('(')
|| result
.chars()
.all(|c| c.is_alphanumeric() || c == '_' || c == '.' || c == '-');
if !is_simple {
let tmp = format!("_t{}", *counter);
*counter += 1;
stmts.push(format!(" double {tmp} = {result};"));
memo.insert(expr, tmp.clone());
return Ok(tmp);
}
memo.insert(expr, result.clone());
Ok(result)
}
fn emit_predicate_inner(
pred: ExprId,
var_map: &HashMap<ExprId, &str>,
pool: &ExprPool,
stmts: &mut Vec<String>,
memo: &mut HashMap<ExprId, String>,
counter: &mut usize,
) -> Result<String, EmitCError> {
use crate::kernel::expr::PredicateKind;
match pool.get(pred) {
ExprData::Predicate { kind, args } => match kind {
PredicateKind::True => Ok("1".to_string()),
PredicateKind::False => Ok("0".to_string()),
PredicateKind::Lt => {
let l = emit_expr_inner(args[0], var_map, pool, stmts, memo, counter)?;
let r = emit_expr_inner(args[1], var_map, pool, stmts, memo, counter)?;
Ok(format!("({l}) < ({r})"))
}
PredicateKind::Le => {
let l = emit_expr_inner(args[0], var_map, pool, stmts, memo, counter)?;
let r = emit_expr_inner(args[1], var_map, pool, stmts, memo, counter)?;
Ok(format!("({l}) <= ({r})"))
}
PredicateKind::Gt => {
let l = emit_expr_inner(args[0], var_map, pool, stmts, memo, counter)?;
let r = emit_expr_inner(args[1], var_map, pool, stmts, memo, counter)?;
Ok(format!("({l}) > ({r})"))
}
PredicateKind::Ge => {
let l = emit_expr_inner(args[0], var_map, pool, stmts, memo, counter)?;
let r = emit_expr_inner(args[1], var_map, pool, stmts, memo, counter)?;
Ok(format!("({l}) >= ({r})"))
}
PredicateKind::Eq => {
let l = emit_expr_inner(args[0], var_map, pool, stmts, memo, counter)?;
let r = emit_expr_inner(args[1], var_map, pool, stmts, memo, counter)?;
Ok(format!("({l}) == ({r})"))
}
PredicateKind::Ne => {
let l = emit_expr_inner(args[0], var_map, pool, stmts, memo, counter)?;
let r = emit_expr_inner(args[1], var_map, pool, stmts, memo, counter)?;
Ok(format!("({l}) != ({r})"))
}
PredicateKind::Not => {
let inner = emit_predicate_inner(args[0], var_map, pool, stmts, memo, counter)?;
Ok(format!("!({inner})"))
}
PredicateKind::And => {
let parts: Result<Vec<_>, _> = args
.iter()
.map(|&a| emit_predicate_inner(a, var_map, pool, stmts, memo, counter))
.collect();
Ok(format!("({})", parts?.join(" && ")))
}
PredicateKind::Or => {
let parts: Result<Vec<_>, _> = args
.iter()
.map(|&a| emit_predicate_inner(a, var_map, pool, stmts, memo, counter))
.collect();
Ok(format!("({})", parts?.join(" || ")))
}
},
_ => Err(EmitCError::UnsupportedNode(
"expected a Predicate node in condition position".to_string(),
)),
}
}
pub fn emit_expr_c(
expr: ExprId,
vars: &[ExprId],
var_names: &[&str],
fn_name: &str,
pool: &ExprPool,
) -> Result<String, EmitCError> {
assert_eq!(
vars.len(),
var_names.len(),
"vars and var_names must have the same length"
);
let mut var_map: HashMap<ExprId, &str> = HashMap::new();
for (&id, &name) in vars.iter().zip(var_names.iter()) {
var_map.insert(id, name);
}
let mut stmts: Vec<String> = Vec::new();
let mut memo: HashMap<ExprId, String> = HashMap::new();
let mut counter = 0usize;
let result = emit_expr_inner(expr, &var_map, pool, &mut stmts, &mut memo, &mut counter)?;
let params: Vec<String> = var_names.iter().map(|n| format!("double {n}")).collect();
let params_str = params.join(", ");
let body = if stmts.is_empty() {
format!(" return {result};\n")
} else {
let mut body = stmts.join("\n");
body.push('\n');
body.push_str(&format!(" return {result};\n"));
body
};
Ok(format!("double {fn_name}({params_str}) {{\n{body}}}\n"))
}
pub fn emit_expr_c_vec(
exprs: &[ExprId],
vars: &[ExprId],
var_names: &[&str],
fn_name: &str,
pool: &ExprPool,
) -> Result<String, EmitCError> {
assert_eq!(
vars.len(),
var_names.len(),
"vars and var_names must have the same length"
);
let mut var_map: HashMap<ExprId, &str> = HashMap::new();
for (&id, &name) in vars.iter().zip(var_names.iter()) {
var_map.insert(id, name);
}
let mut stmts: Vec<String> = Vec::new();
let mut memo: HashMap<ExprId, String> = HashMap::new();
let mut counter = 0usize;
let mut result_exprs: Vec<String> = Vec::with_capacity(exprs.len());
for &e in exprs {
result_exprs.push(emit_expr_inner(
e,
&var_map,
pool,
&mut stmts,
&mut memo,
&mut counter,
)?);
}
let mut params: Vec<String> = var_names.iter().map(|n| format!("double {n}")).collect();
params.push("double *out".to_string());
let params_str = params.join(", ");
let assignments: Vec<String> = result_exprs
.iter()
.enumerate()
.map(|(i, r)| format!(" out[{i}] = {r};"))
.collect();
let body = if stmts.is_empty() {
format!("{}\n", assignments.join("\n"))
} else {
format!("{}\n{}\n", stmts.join("\n"), assignments.join("\n"))
};
Ok(format!("void {fn_name}({params_str}) {{\n{body}}}\n"))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::jit::eval_interp;
use crate::kernel::{Domain, ExprPool};
use std::collections::HashMap;
fn p() -> ExprPool {
ExprPool::new()
}
#[test]
fn horner_linear() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let expr = pool.add(vec![
pool.mul(vec![pool.integer(2_i32), x]),
pool.integer(1_i32),
]);
let h = horner(expr, x, &pool).unwrap();
let mut env = HashMap::new();
env.insert(x, 3.0f64);
let val = eval_interp(h, &env, &pool).unwrap();
assert!((val - 7.0).abs() < 1e-10, "expected 7.0, got {val}");
}
#[test]
fn horner_quadratic() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let x2 = pool.pow(x, pool.integer(2_i32));
let two_x = pool.mul(vec![pool.integer(2_i32), x]);
let one = pool.integer(1_i32);
let expr = pool.add(vec![x2, two_x, one]);
let h = horner(expr, x, &pool).unwrap();
let mut env = HashMap::new();
for v in [-2.0f64, -1.0, 0.0, 1.0, 2.0, 3.0] {
env.insert(x, v);
let expected = (v + 1.0).powi(2);
let actual = eval_interp(h, &env, &pool).unwrap();
assert!(
(actual - expected).abs() < 1e-9,
"v={v}: expected {expected}, got {actual}"
);
}
}
#[test]
fn horner_degree_10_op_count() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let mut expr = pool.integer(1_i32);
for k in 1_i32..=10 {
let xk = pool.pow(x, pool.integer(k));
expr = pool.add(vec![expr, xk]);
}
let h = horner(expr, x, &pool).unwrap();
let muls = count_muls(h, &pool);
assert!(
muls <= 10,
"Horner form should use ≤ 10 multiplications, got {muls}"
);
}
fn count_muls(expr: ExprId, pool: &ExprPool) -> usize {
use crate::kernel::ExprData;
match pool.get(expr) {
ExprData::Mul(args) => 1 + args.iter().map(|&a| count_muls(a, pool)).sum::<usize>(),
ExprData::Add(args) => args.iter().map(|&a| count_muls(a, pool)).sum(),
ExprData::Pow { base, exp } => count_muls(base, pool) + count_muls(exp, pool),
_ => 0,
}
}
#[test]
fn emit_horner_c_quadratic() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let x2 = pool.pow(x, pool.integer(2_i32));
let two_x = pool.mul(vec![pool.integer(2_i32), x]);
let one = pool.integer(1_i32);
let expr = pool.add(vec![x2, two_x, one]);
let code = emit_horner_c(expr, x, "x", "eval_quad", &pool).unwrap();
assert!(code.contains("eval_quad"), "function name not in output");
assert!(code.contains("double"), "return type not in output");
}
#[test]
fn horner_constant() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let five = pool.integer(5_i32);
let h = horner(five, x, &pool).unwrap();
let env = HashMap::new();
let val = eval_interp(h, &env, &pool).unwrap();
assert!((val - 5.0).abs() < 1e-10);
}
#[test]
fn eval_horner_f64_matches_interp() {
let coeffs = [1.0, 2.0, 3.0]; let xs = [-1.0, 0.0, 0.5, 2.0, 10.0];
for &x in &xs {
let scalar = eval_horner_f64(&coeffs, x);
let expected = 1.0 + x * (2.0 + x * 3.0);
assert!((scalar - expected).abs() < 1e-12, "x={x}");
}
}
#[test]
fn eval_horner_f64_batch_matches_scalar() {
let coeffs = [1.0, 2.0, 3.0];
let xs = [-1.0, 0.0, 0.5, 2.0, 10.0, 3.0, 7.0];
let mut out = vec![0.0; xs.len()];
eval_horner_f64_batch(&coeffs, &xs, &mut out);
for (i, &x) in xs.iter().enumerate() {
assert!((out[i] - eval_horner_f64(&coeffs, x)).abs() < 1e-12);
}
}
#[test]
fn emit_expr_c_sin_plus_x_squared() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let sin_x = pool.func("sin", vec![x]);
let x2 = pool.pow(x, pool.integer(2_i32));
let expr = pool.add(vec![sin_x, x2]);
let code = emit_expr_c(expr, &[x], &["x"], "f", &pool).unwrap();
assert!(code.contains("sin("), "expected sin( in:\n{code}");
assert!(
code.contains("double f(double x)"),
"expected signature:\n{code}"
);
assert!(code.contains("return "), "expected return:\n{code}");
}
#[test]
fn emit_expr_c_transcendentals() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let expr = pool.mul(vec![pool.func("cos", vec![x]), pool.func("exp", vec![x])]);
let code = emit_expr_c(expr, &[x], &["x"], "g", &pool).unwrap();
assert!(code.contains("cos("), "expected cos(:\n{code}");
assert!(code.contains("exp("), "expected exp(:\n{code}");
}
#[test]
fn emit_expr_c_multivar() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let x2 = pool.pow(x, pool.integer(2_i32));
let y2 = pool.pow(y, pool.integer(2_i32));
let inner = pool.add(vec![x2, y2]);
let expr = pool.func("sqrt", vec![inner]);
let code = emit_expr_c(expr, &[x, y], &["x", "y"], "norm", &pool).unwrap();
assert!(code.contains("sqrt("), "expected sqrt(:\n{code}");
assert!(code.contains("double x"), "expected x param:\n{code}");
assert!(code.contains("double y"), "expected y param:\n{code}");
}
#[test]
fn emit_expr_c_unsupported_func_errors() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let expr = pool.func("diracdelta", vec![x]);
let err = emit_expr_c(expr, &[x], &["x"], "f", &pool).unwrap_err();
assert!(
matches!(err, EmitCError::UnsupportedFunction(ref n) if n == "diracdelta"),
"unexpected error: {err}"
);
}
#[test]
fn emit_expr_c_missing_var_errors() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let expr = pool.add(vec![x, y]);
let err = emit_expr_c(expr, &[x], &["x"], "f", &pool).unwrap_err();
assert!(
matches!(err, EmitCError::MissingVariable(ref n) if n == "y"),
"unexpected error: {err}"
);
}
#[test]
fn emit_expr_c_vec_two_outputs() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let y = pool.symbol("y", Domain::Real);
let f0 = pool.func("sin", vec![x]);
let f1 = pool.add(vec![pool.pow(x, pool.integer(2_i32)), y]);
let code = emit_expr_c_vec(&[f0, f1], &[x, y], &["x", "y"], "eval_vec", &pool).unwrap();
assert!(
code.contains("double *out"),
"expected out pointer:\n{code}"
);
assert!(code.contains("out[0]"), "expected out[0]:\n{code}");
assert!(code.contains("out[1]"), "expected out[1]:\n{code}");
assert!(code.contains("sin("), "expected sin(:\n{code}");
}
#[test]
fn emit_expr_c_polynomial_still_works() {
let pool = p();
let x = pool.symbol("x", Domain::Real);
let x2 = pool.pow(x, pool.integer(2_i32));
let two_x = pool.mul(vec![pool.integer(2_i32), x]);
let one = pool.integer(1_i32);
let expr = pool.add(vec![x2, two_x, one]);
let code = emit_expr_c(expr, &[x], &["x"], "eval_poly_new", &pool).unwrap();
assert!(code.contains("double eval_poly_new"), "signature:\n{code}");
assert!(code.contains("return "), "return:\n{code}");
}
}