use crate::error::{MathError, Result};
use crate::expr::Expr;
use crate::eval::Context;
use std::ops::{Add, Div, Mul, Neg, Sub};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Dual {
pub val: f64,
pub deriv: f64,
}
impl Dual {
pub fn constant(val: f64) -> Self {
Dual { val, deriv: 0.0 }
}
pub fn var(val: f64) -> Self {
Dual { val, deriv: 1.0 }
}
pub fn new(val: f64, deriv: f64) -> Self {
Dual { val, deriv }
}
pub fn sin(self) -> Self {
Dual {
val: self.val.sin(),
deriv: self.deriv * self.val.cos(),
}
}
pub fn cos(self) -> Self {
Dual {
val: self.val.cos(),
deriv: -self.deriv * self.val.sin(),
}
}
pub fn tan(self) -> Self {
let t = self.val.tan();
Dual {
val: t,
deriv: self.deriv * (1.0 + t * t),
}
}
pub fn exp(self) -> Self {
let e = self.val.exp();
Dual {
val: e,
deriv: self.deriv * e,
}
}
pub fn ln(self) -> Self {
Dual {
val: self.val.ln(),
deriv: self.deriv / self.val,
}
}
pub fn log(self, base: f64) -> Self {
let ln_base = base.ln();
Dual {
val: self.val.ln() / ln_base,
deriv: self.deriv / (self.val * ln_base),
}
}
pub fn sqrt(self) -> Self {
let s = self.val.sqrt();
Dual {
val: s,
deriv: self.deriv / (2.0 * s),
}
}
pub fn powf(self, c: f64) -> Self {
let v = self.val.powf(c);
let d = if self.val != 0.0 {
c * self.val.powf(c - 1.0) * self.deriv
} else if c > 1.0 {
0.0 } else {
f64::INFINITY
};
Dual { val: v, deriv: d }
}
pub fn pow_dual(self, other: Dual) -> Self {
if other.deriv == 0.0 {
self.powf(other.val)
} else {
let ln_a = self.val.ln();
let val = self.val.powf(other.val);
let deriv = val * (other.deriv * ln_a + other.val * self.deriv / self.val);
Dual { val, deriv }
}
}
pub fn asin(self) -> Self {
Dual {
val: self.val.asin(),
deriv: self.deriv / (1.0 - self.val * self.val).sqrt(),
}
}
pub fn acos(self) -> Self {
Dual {
val: self.val.acos(),
deriv: -self.deriv / (1.0 - self.val * self.val).sqrt(),
}
}
pub fn atan(self) -> Self {
Dual {
val: self.val.atan(),
deriv: self.deriv / (1.0 + self.val * self.val),
}
}
pub fn sinh(self) -> Self {
Dual {
val: self.val.sinh(),
deriv: self.deriv * self.val.cosh(),
}
}
pub fn cosh(self) -> Self {
Dual {
val: self.val.cosh(),
deriv: self.deriv * self.val.sinh(),
}
}
pub fn tanh(self) -> Self {
let t = self.val.tanh();
Dual {
val: t,
deriv: self.deriv * (1.0 - t * t),
}
}
pub fn abs(self) -> Self {
Dual {
val: self.val.abs(),
deriv: self.deriv * self.val.signum(),
}
}
}
impl Add for Dual {
type Output = Dual;
fn add(self, other: Dual) -> Dual {
Dual {
val: self.val + other.val,
deriv: self.deriv + other.deriv,
}
}
}
impl Sub for Dual {
type Output = Dual;
fn sub(self, other: Dual) -> Dual {
Dual {
val: self.val - other.val,
deriv: self.deriv - other.deriv,
}
}
}
impl Mul for Dual {
type Output = Dual;
fn mul(self, other: Dual) -> Dual {
Dual {
val: self.val * other.val,
deriv: self.val * other.deriv + self.deriv * other.val,
}
}
}
impl Div for Dual {
type Output = Dual;
fn div(self, other: Dual) -> Dual {
let val = self.val / other.val;
let deriv = (self.deriv * other.val - self.val * other.deriv) / (other.val * other.val);
Dual { val, deriv }
}
}
impl Neg for Dual {
type Output = Dual;
fn neg(self) -> Dual {
Dual {
val: -self.val,
deriv: -self.deriv,
}
}
}
pub fn eval(expr: &Expr, var: &str, x: f64, ctx: &Context) -> Result<Dual> {
match expr {
Expr::Num(n) => Ok(Dual::constant(*n)),
Expr::Var(name) => {
if name == var {
Ok(Dual::var(x))
} else {
let v = ctx
.vars
.get(name)
.ok_or_else(|| MathError::UnknownVariable(name.clone()))?;
Ok(Dual::constant(*v))
}
}
Expr::Neg(e) => Ok(-eval(e, var, x, ctx)?),
Expr::Add(a, b) => Ok(eval(a, var, x, ctx)? + eval(b, var, x, ctx)?),
Expr::Sub(a, b) => Ok(eval(a, var, x, ctx)? - eval(b, var, x, ctx)?),
Expr::Mul(a, b) => Ok(eval(a, var, x, ctx)? * eval(b, var, x, ctx)?),
Expr::Div(a, b) => Ok(eval(a, var, x, ctx)? / eval(b, var, x, ctx)?),
Expr::Pow(a, b) => {
let da = eval(a, var, x, ctx)?;
let db = eval(b, var, x, ctx)?;
Ok(da.pow_dual(db))
}
Expr::Func(name, args) => {
if args.len() != 1 {
let h = 1e-8;
let mut cx0 = ctx.clone();
cx0.set(var, x);
let v0 = crate::eval::eval(expr, &cx0)?;
let mut cx1 = ctx.clone();
cx1.set(var, x + h);
let v1 = crate::eval::eval(expr, &cx1)?;
return Ok(Dual::new(v0, (v1 - v0) / h));
}
let d = eval(&args[0], var, x, ctx)?;
match name.as_str() {
"sin" => Ok(d.sin()),
"cos" => Ok(d.cos()),
"tan" => Ok(d.tan()),
"exp" => Ok(d.exp()),
"ln" | "log" if name == "ln" => Ok(d.ln()),
"log" => Ok(d.log(10.0)),
"log2" => Ok(d.log(2.0)),
"sqrt" => Ok(d.sqrt()),
"asin" => Ok(d.asin()),
"acos" => Ok(d.acos()),
"atan" => Ok(d.atan()),
"sinh" => Ok(d.sinh()),
"cosh" => Ok(d.cosh()),
"tanh" => Ok(d.tanh()),
"abs" => Ok(d.abs()),
_ => {
let h = 1e-8;
let v0 = d.val;
let v1 = {
let mut cx = ctx.clone();
cx.set(var, x + h);
crate::eval::eval(expr, &cx)?
};
let deriv = (v1 - v0) / h;
Ok(Dual::new(v0, deriv))
}
}
}
}
}
pub fn derivative(expr: &Expr, var: &str, x: f64, ctx: &Context) -> Result<Dual> {
eval(expr, var, x, ctx)
}
pub fn gradient(expr: &Expr, point: &Context) -> Result<Vec<(String, f64)>> {
let vars = collect_vars(expr);
let mut result = Vec::with_capacity(vars.len());
for var in &vars {
let x = point
.vars
.get(var)
.ok_or_else(|| MathError::UnknownVariable(var.clone()))?;
let d = eval(expr, var, *x, point)?;
result.push((var.clone(), d.deriv));
}
Ok(result)
}
pub fn jacobian(exprs: &[Expr], point: &Context) -> Result<Vec<Vec<f64>>> {
let mut all_vars: std::collections::HashSet<String> = std::collections::HashSet::new();
for expr in exprs {
for v in collect_vars(expr) {
all_vars.insert(v);
}
}
let mut all_vars: Vec<String> = all_vars.into_iter().collect();
all_vars.sort();
let mut jacobian = Vec::with_capacity(exprs.len());
for expr in exprs {
let mut row = Vec::with_capacity(all_vars.len());
for var in &all_vars {
let x = point
.vars
.get(var)
.ok_or_else(|| MathError::UnknownVariable(var.clone()))?;
let d = eval(expr, var, *x, point)?;
row.push(d.deriv);
}
jacobian.push(row);
}
Ok(jacobian)
}
fn collect_vars(expr: &Expr) -> Vec<String> {
let mut vars = std::collections::HashSet::new();
collect_vars_inner(expr, &mut vars);
let mut v: Vec<_> = vars.into_iter().collect();
v.sort();
v
}
fn collect_vars_inner(expr: &Expr, vars: &mut std::collections::HashSet<String>) {
match expr {
Expr::Var(name) => {
vars.insert(name.clone());
}
Expr::Num(_) => {}
Expr::Neg(e) => collect_vars_inner(e, vars),
Expr::Add(a, b) | Expr::Sub(a, b) | Expr::Mul(a, b) | Expr::Div(a, b) | Expr::Pow(a, b) => {
collect_vars_inner(a, vars);
collect_vars_inner(b, vars);
}
Expr::Func(_, args) => {
for a in args {
collect_vars_inner(a, vars);
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::parser::Parser;
fn close(a: f64, b: f64, tol: f64) -> bool {
(a - b).abs() < tol
}
#[test]
fn dual_arithmetic() {
let a = Dual::new(3.0, 1.0);
let b = Dual::new(2.0, 0.0);
let c = a + b;
assert!(close(c.val, 5.0, 1e-10) && close(c.deriv, 1.0, 1e-10));
let d = a * b;
assert!(close(d.val, 6.0, 1e-10) && close(d.deriv, 2.0, 1e-10));
let e = a / b;
assert!(close(e.val, 1.5, 1e-10) && close(e.deriv, 0.5, 1e-10));
}
#[test]
fn dual_sin() {
let d = Dual::var(0.0).sin();
assert!(close(d.val, 0.0, 1e-10));
assert!(close(d.deriv, 1.0, 1e-10)); }
#[test]
fn dual_exp() {
let d = Dual::var(0.0).exp();
assert!(close(d.val, 1.0, 1e-10));
assert!(close(d.deriv, 1.0, 1e-10)); }
#[test]
fn dual_powf() {
let d = Dual::var(3.0).powf(2.0);
assert!(close(d.val, 9.0, 1e-10));
assert!(close(d.deriv, 6.0, 1e-10)); }
#[test]
fn dual_log() {
let d = Dual::var(1.0).ln();
assert!(close(d.val, 0.0, 1e-10));
assert!(close(d.deriv, 1.0, 1e-10)); }
#[test]
fn dual_sqrt() {
let d = Dual::var(4.0).sqrt();
assert!(close(d.val, 2.0, 1e-10));
assert!(close(d.deriv, 0.25, 1e-10)); }
#[test]
fn dual_tan() {
let d = Dual::var(0.0).tan();
assert!(close(d.val, 0.0, 1e-10));
assert!(close(d.deriv, 1.0, 1e-10)); }
#[test]
fn dual_atan() {
let d = Dual::var(0.0).atan();
assert!(close(d.val, 0.0, 1e-10));
assert!(close(d.deriv, 1.0, 1e-10)); }
#[test]
fn dual_sinh_cosh_tanh() {
let s = Dual::var(0.0).sinh();
assert!(close(s.val, 0.0, 1e-10) && close(s.deriv, 1.0, 1e-10));
let c = Dual::var(0.0).cosh();
assert!(close(c.val, 1.0, 1e-10) && close(c.deriv, 0.0, 1e-10));
let t = Dual::var(0.0).tanh();
assert!(close(t.val, 0.0, 1e-10) && close(t.deriv, 1.0, 1e-10));
}
#[test]
fn autodiff_polynomial() {
let expr = Parser::parse("x^3 + 2*x^2 - x + 5").unwrap();
let ctx = Context::standard();
for x in [-2.0, -1.0, 0.0, 1.0, 3.5] {
let d = derivative(&expr, "x", x, &ctx).unwrap();
let expected_val = x * x * x + 2.0 * x * x - x + 5.0;
let expected_deriv = 3.0 * x * x + 4.0 * x - 1.0;
assert!(close(d.val, expected_val, 1e-10), "val at x={}", x);
assert!(close(d.deriv, expected_deriv, 1e-10), "deriv at x={}", x);
}
}
#[test]
fn autodiff_trig_composition() {
let expr = Parser::parse("sin(x^2)").unwrap();
let ctx = Context::standard();
for x in [0.0, 0.5, 1.0, 2.0] {
let d = derivative(&expr, "x", x, &ctx).unwrap();
let expected_val = (x * x).sin();
let expected_deriv = 2.0 * x * (x * x).cos();
assert!(close(d.val, expected_val, 1e-10), "val at x={}", x);
assert!(close(d.deriv, expected_deriv, 1e-10), "deriv at x={}", x);
}
}
#[test]
fn autodiff_exp_composition() {
let expr = Parser::parse("exp(-(x^2))").unwrap();
let ctx = Context::standard();
for x in [0.0, 0.5, 1.0, 2.0] {
let d = derivative(&expr, "x", x, &ctx).unwrap();
let expected_val = (-x * x).exp();
let expected_deriv = -2.0 * x * (-x * x).exp();
assert!(close(d.val, expected_val, 1e-10), "val at x={}", x);
assert!(close(d.deriv, expected_deriv, 1e-10), "deriv at x={}", x);
}
}
#[test]
fn autodiff_log_sqrt() {
let expr = Parser::parse("ln(sqrt(x))").unwrap();
let ctx = Context::standard();
for x in [1.0, 4.0, 9.0, 100.0] {
let d = derivative(&expr, "x", x, &ctx).unwrap();
let expected_val = x.sqrt().ln();
let expected_deriv = 1.0 / (2.0 * x);
assert!(close(d.val, expected_val, 1e-10), "val at x={}", x);
assert!(close(d.deriv, expected_deriv, 1e-10), "deriv at x={}", x);
}
}
#[test]
fn autodiff_product_rule() {
let expr = Parser::parse("x * sin(x)").unwrap();
let ctx = Context::standard();
for x in [0.0, 1.0, 2.0] {
let d = derivative(&expr, "x", x, &ctx).unwrap();
let expected_val = x * x.sin();
let expected_deriv = x.sin() + x * x.cos();
assert!(close(d.val, expected_val, 1e-10), "val at x={}", x);
assert!(close(d.deriv, expected_deriv, 1e-10), "deriv at x={}", x);
}
}
#[test]
fn autodiff_quotient_rule() {
let expr = Parser::parse("sin(x) / x").unwrap();
let ctx = Context::standard();
for x in [1.0, 2.0, 5.0] {
let d = derivative(&expr, "x", x, &ctx).unwrap();
let expected_val = x.sin() / x;
let expected_deriv = (x * x.cos() - x.sin()) / (x * x);
assert!(close(d.val, expected_val, 1e-10), "val at x={}", x);
assert!(close(d.deriv, expected_deriv, 1e-10), "deriv at x={}", x);
}
}
#[test]
fn autodiff_chain_rule_deep() {
let expr = Parser::parse("cos(sin(x^2 + 1))").unwrap();
let ctx = Context::standard();
for x in [0.5, 1.0, 1.5] {
let d = derivative(&expr, "x", x, &ctx).unwrap();
let inner = x * x + 1.0;
let expected_val = inner.sin().cos();
let expected_deriv = -inner.sin().sin() * inner.cos() * 2.0 * x;
assert!(close(d.val, expected_val, 1e-10), "val at x={}", x);
assert!(close(d.deriv, expected_deriv, 1e-10), "deriv at x={}", x);
}
}
#[test]
fn autodiff_gradient() {
let expr = Parser::parse("x^2 + y^3").unwrap();
let mut ctx = Context::standard();
ctx.set("x", 2.0);
ctx.set("y", 3.0);
let grad = gradient(&expr, &ctx).unwrap();
let x_grad = grad.iter().find(|(name, _)| name == "x").unwrap().1;
let y_grad = grad.iter().find(|(name, _)| name == "y").unwrap().1;
assert!(close(x_grad, 4.0, 1e-10)); assert!(close(y_grad, 27.0, 1e-10)); }
#[test]
fn autodiff_jacobian() {
let f1 = Parser::parse("x^2 + y").unwrap();
let f2 = Parser::parse("x * y^2").unwrap();
let mut ctx = Context::standard();
ctx.set("x", 2.0);
ctx.set("y", 3.0);
let jac = jacobian(&[f1, f2], &ctx).unwrap();
assert!(close(jac[0][0], 4.0, 1e-10));
assert!(close(jac[0][1], 1.0, 1e-10));
assert!(close(jac[1][0], 9.0, 1e-10));
assert!(close(jac[1][1], 12.0, 1e-10));
}
#[test]
fn autodiff_constant() {
let expr = Parser::parse("42").unwrap();
let ctx = Context::standard();
let d = derivative(&expr, "x", 1.0, &ctx).unwrap();
assert!(close(d.val, 42.0, 1e-10));
assert!(close(d.deriv, 0.0, 1e-10));
}
#[test]
fn autodiff_var_not_in_expr() {
let expr = Parser::parse("y^2").unwrap();
let mut ctx = Context::standard();
ctx.set("y", 5.0);
let d = derivative(&expr, "x", 1.0, &ctx).unwrap();
assert!(close(d.val, 25.0, 1e-10));
assert!(close(d.deriv, 0.0, 1e-10));
}
#[test]
fn autodiff_powf_general() {
let expr = Parser::parse("x^x").unwrap();
let ctx = Context::standard();
for x in [1.0, 2.0, 3.0] {
let d = derivative(&expr, "x", x, &ctx).unwrap();
let expected_val = x.powf(x);
let expected_deriv = x.powf(x) * (x.ln() + 1.0);
assert!(close(d.val, expected_val, 1e-8), "val at x={}", x);
assert!(close(d.deriv, expected_deriv, 1e-8), "deriv at x={}", x);
}
}
}