use rustc_hash::{FxHashMap, FxHashSet};
use smallvec::SmallVec;
use crate::base::arena::Arena;
use crate::base::libfn::LibFn;
use crate::base::node::{ExprId, ExprNode, SymbolId};
use crate::transforms::eval::apply_named;
pub(crate) fn diff(arena: &mut Arena, expr: ExprId, var: ExprId) -> ExprId {
let var_sym = match arena.node(var) {
ExprNode::Symbol(sid) => *sid,
_ => return arena.zero,
};
let post_order = crate::base::walk::post_order_ids(arena, expr);
let mut cache: FxHashMap<ExprId, ExprId> = FxHashMap::default();
for &id in &post_order {
let deriv = diff_node(arena, id, var_sym, &cache);
cache.insert(id, deriv);
}
cache.get(&expr).copied().unwrap_or(arena.zero)
}
pub(crate) fn diff_with_deps(
arena: &mut Arena,
expr: ExprId,
var: ExprId,
deps: &FxHashSet<ExprId>,
) -> ExprId {
let var_sym = match arena.node(var) {
ExprNode::Symbol(sid) => *sid,
_ => return arena.zero,
};
let post_order = crate::base::walk::post_order_ids(arena, expr);
let mut cache: FxHashMap<ExprId, ExprId> = FxHashMap::default();
for &dep_id in deps {
let formal = arena.intern(ExprNode::Derivative(dep_id, var));
cache.insert(dep_id, formal);
}
for &id in &post_order {
if cache.contains_key(&id) {
continue; }
let deriv = diff_node(arena, id, var_sym, &cache);
cache.insert(id, deriv);
}
cache.get(&expr).copied().unwrap_or(arena.zero)
}
fn diff_node(
arena: &mut Arena,
id: ExprId,
var: SymbolId,
cache: &FxHashMap<ExprId, ExprId>,
) -> ExprId {
let node = arena.node(id).clone();
match node {
ExprNode::Num(_) => arena.zero,
ExprNode::Symbol(sid) => {
if sid == var {
arena.one
} else {
arena.zero
}
}
ExprNode::Pi
| ExprNode::E
| ExprNode::ImaginaryUnit
| ExprNode::EulerGamma
| ExprNode::Catalan
| ExprNode::GoldenRatio
| ExprNode::PhysicalConstant(_, _)
| ExprNode::Infinity
| ExprNode::NegInfinity
| ExprNode::ComplexInfinity
| ExprNode::NaN
| ExprNode::BoolTrue
| ExprNode::BoolFalse => arena.zero,
ExprNode::Gt(..)
| ExprNode::Ge(..)
| ExprNode::Eq_(..)
| ExprNode::Ne(..)
| ExprNode::And(_)
| ExprNode::Or(_)
| ExprNode::Not(_) => arena.zero,
ExprNode::EmptySet
| ExprNode::UniversalSet
| ExprNode::Interval(..)
| ExprNode::FiniteSet(_)
| ExprNode::SetUnion(_)
| ExprNode::SetIntersection(_)
| ExprNode::SetComplement(..) => arena.zero,
ExprNode::Piecewise(ref pairs) => {
let pairs = pairs.clone();
let mut new_pairs = SmallVec::new();
for &(val, cond) in &pairs {
let dval = get_deriv(cache, val, arena);
new_pairs.push((dval, cond));
}
arena.intern(ExprNode::Piecewise(new_pairs))
}
ExprNode::Add(ref children) => {
let derivs: SmallVec<[ExprId; 6]> = children
.iter()
.map(|&child| get_deriv(cache, child, arena))
.collect();
arena.add(&derivs)
}
ExprNode::Mul(ref children) => {
let n = children.len();
if n == 0 {
return arena.zero;
}
let children = children.clone();
let mut sum_terms: SmallVec<[ExprId; 6]> = SmallVec::new();
for i in 0..n {
let di = get_deriv(cache, children[i], arena);
if arena.is_zero_structural(di) {
continue;
}
let mut factors: SmallVec<[ExprId; 6]> = SmallVec::new();
for (j, &child) in children.iter().enumerate() {
if j == i {
factors.push(di);
} else {
factors.push(child);
}
}
let term = arena.mul(&factors);
sum_terms.push(term);
}
if sum_terms.is_empty() {
arena.zero
} else {
arena.add(&sum_terms)
}
}
ExprNode::Pow(base, exp) => {
let dbase = get_deriv(cache, base, arena);
let dexp = get_deriv(cache, exp, arena);
let base_is_const = arena.is_zero_structural(dbase);
let exp_is_const = arena.is_zero_structural(dexp);
if base_is_const && exp_is_const {
arena.zero
} else if exp_is_const {
let n_minus_1 = arena.sub(exp, arena.one);
let pow_part = arena.pow(base, n_minus_1);
arena.mul(&[exp, pow_part, dbase])
} else if base_is_const {
let ln_base = arena.ln(base);
let pow_part = arena.pow(base, exp);
arena.mul(&[pow_part, ln_base, dexp])
} else {
let pow_part = arena.pow(base, exp);
let ln_f = arena.ln(base);
let term1 = arena.mul(&[dexp, ln_f]);
let neg_one = arena.neg_one;
let f_inv = arena.pow(base, neg_one);
let term2 = arena.mul(&[exp, dbase, f_inv]);
let inner = arena.add(&[term1, term2]);
arena.mul(&[pow_part, inner])
}
}
ExprNode::Neg(inner) => {
let di = get_deriv(cache, inner, arena);
arena.neg(di)
}
ExprNode::Sin(inner) => {
let di = get_deriv(cache, inner, arena);
if arena.is_zero_structural(di) {
return arena.zero;
}
let cos_f = arena.cos(inner);
arena.mul(&[cos_f, di])
}
ExprNode::Cos(inner) => {
let di = get_deriv(cache, inner, arena);
if arena.is_zero_structural(di) {
return arena.zero;
}
let sin_f = arena.sin(inner);
let neg_sin = arena.neg(sin_f);
arena.mul(&[neg_sin, di])
}
ExprNode::Tan(inner) => {
let di = get_deriv(cache, inner, arena);
if arena.is_zero_structural(di) {
return arena.zero;
}
let tan_f = arena.tan(inner);
let two = arena.int(2);
let tan_sq = arena.pow(tan_f, two);
let one = arena.one;
let one_plus_tan_sq = arena.add(&[one, tan_sq]);
arena.mul(&[one_plus_tan_sq, di])
}
ExprNode::Exp(inner) => {
let di = get_deriv(cache, inner, arena);
if arena.is_zero_structural(di) {
return arena.zero;
}
let exp_f = arena.exp(inner);
arena.mul(&[exp_f, di])
}
ExprNode::Ln(inner) => {
let di = get_deriv(cache, inner, arena);
if arena.is_zero_structural(di) {
return arena.zero;
}
arena.div(di, inner)
}
ExprNode::Abs(inner) => {
let inner_diff = cache.get(&inner).copied().unwrap_or(arena.zero);
let sign_f = arena.sign(inner);
arena.mul(&[sign_f, inner_diff])
}
ExprNode::Sign(_) => arena.zero,
ExprNode::Heaviside(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let delta = arena.intern(ExprNode::DiracDelta(inner));
arena.mul(&[delta, df])
}
ExprNode::DiracDelta(_) => {
let v = var_expr(arena, var);
arena.intern(ExprNode::Derivative(id, v))
}
ExprNode::Asin(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let one = arena.one;
let two = arena.int(2);
let f_sq = arena.pow(inner, two);
let one_minus_f_sq = arena.sub(one, f_sq);
let sqrt_denom = arena.sqrt(one_minus_f_sq);
arena.div(df, sqrt_denom)
}
ExprNode::Acos(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let one = arena.one;
let two = arena.int(2);
let f_sq = arena.pow(inner, two);
let one_minus_f_sq = arena.sub(one, f_sq);
let sqrt_denom = arena.sqrt(one_minus_f_sq);
let frac = arena.div(df, sqrt_denom);
arena.neg(frac)
}
ExprNode::Atan(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let one = arena.one;
let two = arena.int(2);
let f_sq = arena.pow(inner, two);
let one_plus_f_sq = arena.add(&[one, f_sq]);
arena.div(df, one_plus_f_sq)
}
ExprNode::Atan2(y_id, x_id) => {
let dy = get_deriv(cache, y_id, arena);
let dx = get_deriv(cache, x_id, arena);
let both_zero = arena.is_zero_structural(dy) && arena.is_zero_structural(dx);
if both_zero {
return arena.zero;
}
let two = arena.int(2);
let x_sq = arena.pow(x_id, two);
let y_sq = arena.pow(y_id, two);
let denom = arena.add(&[x_sq, y_sq]);
let x_dy = arena.mul(&[x_id, dy]);
let y_dx = arena.mul(&[y_id, dx]);
let numer = arena.sub(x_dy, y_dx);
arena.div(numer, denom)
}
ExprNode::Sinh(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let cosh_f = arena.intern(ExprNode::Cosh(inner));
arena.mul(&[cosh_f, df])
}
ExprNode::Cosh(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let sinh_f = arena.intern(ExprNode::Sinh(inner));
arena.mul(&[sinh_f, df])
}
ExprNode::Tanh(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let tanh_f = arena.intern(ExprNode::Tanh(inner));
let two = arena.int(2);
let tanh_sq = arena.pow(tanh_f, two);
let one = arena.one;
let one_minus_tanh_sq = arena.sub(one, tanh_sq);
arena.mul(&[one_minus_tanh_sq, df])
}
ExprNode::Asinh(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let one = arena.one;
let two = arena.int(2);
let f_sq = arena.pow(inner, two);
let f_sq_plus_1 = arena.add(&[f_sq, one]);
let sqrt_denom = arena.sqrt(f_sq_plus_1);
arena.div(df, sqrt_denom)
}
ExprNode::Acosh(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let one = arena.one;
let two = arena.int(2);
let f_sq = arena.pow(inner, two);
let f_sq_minus_1 = arena.sub(f_sq, one);
let sqrt_denom = arena.sqrt(f_sq_minus_1);
arena.div(df, sqrt_denom)
}
ExprNode::Atanh(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let one = arena.one;
let two = arena.int(2);
let f_sq = arena.pow(inner, two);
let one_minus_f_sq = arena.sub(one, f_sq);
arena.div(df, one_minus_f_sq)
}
ExprNode::Gamma(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let gamma_f = arena.gamma(inner);
let digamma_f = arena.digamma(inner);
arena.mul(&[gamma_f, digamma_f, df])
}
ExprNode::LogGamma(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let digamma_f = arena.digamma(inner);
arena.mul(&[digamma_f, df])
}
ExprNode::Digamma(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let one = arena.one;
let trigamma = arena.polygamma(one, inner);
arena.mul(&[trigamma, df])
}
ExprNode::Polygamma(n, inner) => {
let dn = get_deriv(cache, n, arena);
let df = get_deriv(cache, inner, arena);
if !arena.is_zero_structural(dn) {
let v = var_expr(arena, var);
return arena.intern(ExprNode::Derivative(id, v));
}
if arena.is_zero_structural(df) {
return arena.zero;
}
let one = arena.one;
let n_plus_1 = arena.add(&[n, one]);
let next = arena.polygamma(n_plus_1, inner);
arena.mul(&[next, df])
}
ExprNode::Re(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
if !var_is_real(arena, var) {
let v = var_expr(arena, var);
return arena.intern(ExprNode::Derivative(id, v));
}
arena.re(df)
}
ExprNode::Im(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
if !var_is_real(arena, var) {
let v = var_expr(arena, var);
return arena.intern(ExprNode::Derivative(id, v));
}
arena.im(df)
}
ExprNode::Conjugate(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
if !var_is_real(arena, var) {
let v = var_expr(arena, var);
return arena.intern(ExprNode::Derivative(id, v));
}
arena.conjugate(df)
}
ExprNode::Arg(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
if !var_is_real(arena, var) {
let v = var_expr(arena, var);
return arena.intern(ExprNode::Derivative(id, v));
}
let ratio = arena.div(df, inner);
arena.im(ratio)
}
ExprNode::Si(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let sin_f = arena.sin(inner);
let ratio = arena.div(sin_f, inner);
arena.mul(&[ratio, df])
}
ExprNode::Ci(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let cos_f = arena.cos(inner);
let ratio = arena.div(cos_f, inner);
arena.mul(&[ratio, df])
}
ExprNode::Ei(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let exp_f = arena.exp(inner);
let ratio = arena.div(exp_f, inner);
arena.mul(&[ratio, df])
}
ExprNode::Li(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let ln_f = arena.ln(inner);
arena.div(df, ln_f)
}
ExprNode::Zeta(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let v = var_expr(arena, var);
arena.intern(ExprNode::Derivative(id, v))
}
ExprNode::KroneckerDelta(_, _) => arena.zero,
ExprNode::Erf(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let two = arena.int(2);
let pi = arena.pi;
let sqrt_pi = arena.sqrt(pi);
let coeff = arena.div(two, sqrt_pi);
let two2 = arena.int(2);
let f_sq = arena.pow(inner, two2);
let neg_f_sq = arena.neg(f_sq);
let exp_neg_f_sq = arena.exp(neg_f_sq);
arena.mul(&[coeff, exp_neg_f_sq, df])
}
ExprNode::Erfc(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let two = arena.int(2);
let pi = arena.pi;
let sqrt_pi = arena.sqrt(pi);
let coeff = arena.div(two, sqrt_pi);
let neg_coeff = arena.neg(coeff);
let two2 = arena.int(2);
let f_sq = arena.pow(inner, two2);
let neg_f_sq = arena.neg(f_sq);
let exp_neg_f_sq = arena.exp(neg_f_sq);
arena.mul(&[neg_coeff, exp_neg_f_sq, df])
}
ExprNode::LambertW(inner) => {
let df = get_deriv(cache, inner, arena);
if arena.is_zero_structural(df) {
return arena.zero;
}
let w_f = arena.lambertw(inner);
let one = arena.one;
let one_plus_w = arena.add(&[one, w_f]);
let denom = arena.mul(&[inner, one_plus_w]);
let frac = arena.div(w_f, denom);
arena.mul(&[frac, df])
}
ExprNode::Beta(_, _) => {
let v = var_expr(arena, var);
arena.intern(ExprNode::Derivative(id, v))
}
ExprNode::Factorial(_) => {
let v = var_expr(arena, var);
arena.intern(ExprNode::Derivative(id, v))
}
ExprNode::Binomial(_, _) => {
let v = var_expr(arena, var);
arena.intern(ExprNode::Derivative(id, v))
}
ExprNode::Apply(func_sym, ref args) => {
let args_clone = args.clone();
if let Some(f) = arena.lib_fn(func_sym)
&& f.arity().accepts(args_clone.len())
&& let Some(result) = diff_lib_fn(arena, id, f, &args_clone, cache)
{
return result;
}
let mut terms: SmallVec<[ExprId; 4]> = SmallVec::new();
for &arg in &args_clone {
let d_arg = get_deriv(cache, arg, arena);
if arena.is_zero_structural(d_arg) {
continue;
}
let partial = arena.intern(ExprNode::Derivative(id, arg));
terms.push(arena.mul(&[partial, d_arg]));
}
if terms.is_empty() {
arena.zero
} else if terms.len() == 1 {
terms[0]
} else {
arena.add(&terms)
}
}
ExprNode::Derivative(_, _) => {
let v = var_expr(arena, var);
arena.intern(ExprNode::Derivative(id, v))
}
ExprNode::Integral(body, int_var) => {
if let ExprNode::Symbol(int_sym) = arena.node(int_var)
&& *int_sym == var
{
return body;
}
let v = var_expr(arena, var);
arena.intern(ExprNode::Derivative(id, v))
}
ExprNode::DefiniteIntegral(body, int_var, lo, hi) => {
let var_is_bound = matches!(arena.node(int_var), ExprNode::Symbol(s) if *s == var);
let mut terms: SmallVec<[ExprId; 3]> = SmallVec::new();
let d_hi = get_deriv(cache, hi, arena);
if !arena.is_zero_structural(d_hi) {
let f_hi = arena.subs_structural(body, int_var, hi);
terms.push(arena.mul(&[f_hi, d_hi]));
}
let d_lo = get_deriv(cache, lo, arena);
if !arena.is_zero_structural(d_lo) {
let f_lo = arena.subs_structural(body, int_var, lo);
let t = arena.mul(&[f_lo, d_lo]);
terms.push(arena.neg(t));
}
if !var_is_bound {
let d_body = get_deriv(cache, body, arena);
if !arena.is_zero_structural(d_body) {
terms.push(arena.definite_integral(d_body, int_var, lo, hi));
}
}
match terms.len() {
0 => arena.zero,
1 => terms[0],
_ => arena.add(&terms),
}
}
ExprNode::Floor(_) | ExprNode::Ceiling(_) => arena.zero,
ExprNode::Min(_) | ExprNode::Max(_) => {
let v = var_expr(arena, var);
arena.intern(ExprNode::Derivative(id, v))
}
ExprNode::Sum(body, sum_var, lo, hi) => {
if let ExprNode::Symbol(sum_sym) = arena.node(sum_var)
&& *sum_sym == var
{
let v = var_expr(arena, var);
return arena.intern(ExprNode::Derivative(id, v));
}
let dbody = get_deriv(cache, body, arena);
arena.intern(ExprNode::Sum(dbody, sum_var, lo, hi))
}
ExprNode::Product_(_, _, _, _) => {
let v = var_expr(arena, var);
arena.intern(ExprNode::Derivative(id, v))
}
ExprNode::RootSum(poly, body, sumvar) => {
if let ExprNode::Symbol(sum_sym) = arena.node(sumvar)
&& *sum_sym == var
{
let v = var_expr(arena, var);
arena.intern(ExprNode::Derivative(id, v))
} else {
let dbody = get_deriv(cache, body, arena);
tracing::trace!("diff: RootSum — differentiating body w.r.t. outer variable");
arena.intern(ExprNode::RootSum(poly, dbody, sumvar))
}
}
ExprNode::Limit(_, _, _)
| ExprNode::Series(_, _, _, _)
| ExprNode::LaplaceTransform(_, _, _)
| ExprNode::InverseLaplaceTransform(_, _, _)
| ExprNode::Residue(_, _, _)
| ExprNode::RootOf(_, _)
| ExprNode::DSolve(_, _, _)
| ExprNode::ConditionSet(_, _) => {
let v = var_expr(arena, var);
let free = crate::base::walk::free_symbols(arena, id);
if !free.contains(&v) {
return arena.zero;
}
arena.intern(ExprNode::Derivative(id, v))
}
}
}
#[inline]
fn get_deriv(cache: &FxHashMap<ExprId, ExprId>, id: ExprId, arena: &Arena) -> ExprId {
cache.get(&id).copied().unwrap_or(arena.zero)
}
fn var_expr(arena: &mut Arena, var: SymbolId) -> ExprId {
arena.intern(ExprNode::Symbol(var))
}
fn var_is_real(arena: &mut Arena, var: SymbolId) -> bool {
let v = var_expr(arena, var);
crate::base::complex::is_real(arena, v) == Some(true)
}
fn diff_lib_fn(
arena: &mut Arena,
id: ExprId,
f: LibFn,
args: &[ExprId],
cache: &FxHashMap<ExprId, ExprId>,
) -> Option<ExprId> {
match f {
LibFn::BesselJ
| LibFn::BesselY
| LibFn::BesselI
| LibFn::BesselK
| LibFn::Legendre
| LibFn::ChebyshevT
| LibFn::ChebyshevU
| LibFn::Hermite
| LibFn::Laguerre => diff_known_apply(arena, f, args, cache),
LibFn::Erfi
| LibFn::ErfInv
| LibFn::ErfcInv
| LibFn::ExpInt
| LibFn::Shi
| LibFn::Chi
| LibFn::FresnelS
| LibFn::FresnelC
| LibFn::LowerGamma
| LibFn::UpperGamma
| LibFn::PolyLog
| LibFn::DirichletEta
| LibFn::AiryAi
| LibFn::AiryBi
| LibFn::AiryAiPrime
| LibFn::AiryBiPrime
| LibFn::EllipticK
| LibFn::EllipticE
| LibFn::EllipticF
| LibFn::EllipticPi
| LibFn::Gegenbauer
| LibFn::Jacobi
| LibFn::AssocLegendre
| LibFn::AssocLaguerre
| LibFn::BetaInc
| LibFn::BetaIncRegularized => diff_special_09(arena, id, f, args, cache),
LibFn::Factorial2
| LibFn::Subfactorial
| LibFn::RisingFactorial
| LibFn::FallingFactorial
| LibFn::Fibonacci
| LibFn::Lucas
| LibFn::Bernoulli
| LibFn::Harmonic
| LibFn::Catalan
| LibFn::Bell
| LibFn::EulerNumber
| LibFn::Stirling1
| LibFn::Stirling2
| LibFn::PartitionCount
| LibFn::LambertW => None,
}
}
fn diff_known_apply(
arena: &mut Arena,
f: LibFn,
args: &[ExprId],
cache: &FxHashMap<ExprId, ExprId>,
) -> Option<ExprId> {
let [param, x] = args else {
return None;
};
let (param, x) = (*param, *x);
let dparam = get_deriv(cache, param, arena);
if !arena.is_zero_structural(dparam) {
return None;
}
let dx = get_deriv(cache, x, arena);
if arena.is_zero_structural(dx) {
return Some(arena.zero);
}
let one = arena.one;
let two = arena.int(2);
let half = arena.rational(1, 2);
let p_minus_1 = arena.sub(param, one);
let p_plus_1 = arena.add(&[param, one]);
let outer = match f {
LibFn::BesselJ => {
let a = arena.besselj(p_minus_1, x);
let b = arena.besselj(p_plus_1, x);
let d = arena.sub(a, b);
arena.mul(&[half, d])
}
LibFn::BesselY => {
let a = arena.bessely(p_minus_1, x);
let b = arena.bessely(p_plus_1, x);
let d = arena.sub(a, b);
arena.mul(&[half, d])
}
LibFn::BesselI => {
let a = arena.besseli(p_minus_1, x);
let b = arena.besseli(p_plus_1, x);
let s = arena.add(&[a, b]);
arena.mul(&[half, s])
}
LibFn::BesselK => {
let a = arena.besselk(p_minus_1, x);
let b = arena.besselk(p_plus_1, x);
let s = arena.add(&[a, b]);
let neg_half = arena.rational(-1, 2);
arena.mul(&[neg_half, s])
}
LibFn::Legendre => {
let pn = arena.legendre(param, x);
let pn_1 = arena.legendre(p_minus_1, x);
let x_pn = arena.mul(&[x, pn]);
let numer_inner = arena.sub(x_pn, pn_1);
let numer = arena.mul(&[param, numer_inner]);
let x2 = arena.pow(x, two);
let denom = arena.sub(x2, one);
arena.div(numer, denom)
}
LibFn::ChebyshevT => {
let u = arena.chebyshev_u(p_minus_1, x);
arena.mul(&[param, u])
}
LibFn::ChebyshevU => {
let t = arena.chebyshev_t(p_plus_1, x);
let un = arena.chebyshev_u(param, x);
let a = arena.mul(&[p_plus_1, t]);
let b = arena.mul(&[x, un]);
let numer = arena.sub(a, b);
let x2 = arena.pow(x, two);
let denom = arena.sub(x2, one);
arena.div(numer, denom)
}
LibFn::Hermite => {
let h = arena.hermite(p_minus_1, x);
arena.mul(&[two, param, h])
}
LibFn::Laguerre => {
let ln_ = arena.laguerre(param, x);
let ln_1 = arena.laguerre(p_minus_1, x);
let d = arena.sub(ln_, ln_1);
let numer = arena.mul(&[param, d]);
arena.div(numer, x)
}
_ => return None,
};
Some(arena.mul(&[outer, dx]))
}
fn two_over_sqrt_pi(arena: &mut Arena) -> ExprId {
let two = arena.int(2);
let sqrt_pi = arena.sqrt(arena.pi);
arena.div(two, sqrt_pi)
}
fn half_pi_sq(arena: &mut Arena, f: ExprId) -> ExprId {
let two = arena.int(2);
let f2 = arena.pow(f, two);
let half = arena.rational(1, 2);
arena.mul(&[half, arena.pi, f2])
}
fn diff_special_09(
arena: &mut Arena,
id: ExprId,
name: LibFn,
args: &[ExprId],
cache: &FxHashMap<ExprId, ExprId>,
) -> Option<ExprId> {
let formal = |arena: &mut Arena, arg: ExprId| -> ExprId {
let partial = arena.intern(ExprNode::Derivative(id, arg));
let d = get_deriv(cache, arg, arena);
arena.mul(&[partial, d])
};
let unary_rule = |arena: &mut Arena, f: ExprId| -> Option<ExprId> {
Some(match name {
LibFn::Erfi => {
let two = arena.int(2);
let f2 = arena.pow(f, two);
let e = arena.exp(f2);
let c = two_over_sqrt_pi(arena);
arena.mul(&[c, e])
}
LibFn::ErfInv | LibFn::ErfcInv => {
let two = arena.int(2);
let w = apply_named(arena, name, &[f]);
let w2 = arena.pow(w, two);
let e = arena.exp(w2);
let sqrt_pi = arena.sqrt(arena.pi);
let half = arena.rational(if name == LibFn::ErfInv { 1 } else { -1 }, 2);
arena.mul(&[half, sqrt_pi, e])
}
LibFn::Shi => {
let s = arena.sinh(f);
arena.div(s, f)
}
LibFn::Chi => {
let c = arena.cosh(f);
arena.div(c, f)
}
LibFn::FresnelS => {
let a = half_pi_sq(arena, f);
arena.sin(a)
}
LibFn::FresnelC => {
let a = half_pi_sq(arena, f);
arena.cos(a)
}
LibFn::AiryAi => apply_named(arena, LibFn::AiryAiPrime, &[f]),
LibFn::AiryBi => apply_named(arena, LibFn::AiryBiPrime, &[f]),
LibFn::AiryAiPrime => {
let ai = apply_named(arena, LibFn::AiryAi, &[f]);
arena.mul(&[f, ai])
}
LibFn::AiryBiPrime => {
let bi = apply_named(arena, LibFn::AiryBi, &[f]);
arena.mul(&[f, bi])
}
LibFn::EllipticK => {
let k = apply_named(arena, LibFn::EllipticK, &[f]);
let e = apply_named(arena, LibFn::EllipticE, &[f]);
let one_minus_m = arena.sub(arena.one, f);
let t = arena.mul(&[one_minus_m, k]);
let numer = arena.sub(e, t);
let two = arena.int(2);
let denom = arena.mul(&[two, f, one_minus_m]);
arena.div(numer, denom)
}
LibFn::EllipticE => {
let k = apply_named(arena, LibFn::EllipticK, &[f]);
let e = apply_named(arena, LibFn::EllipticE, &[f]);
let numer = arena.sub(e, k);
let two = arena.int(2);
let denom = arena.mul(&[two, f]);
arena.div(numer, denom)
}
_ => return None,
})
};
match (name, args.len()) {
(
LibFn::Erfi
| LibFn::ErfInv
| LibFn::ErfcInv
| LibFn::Shi
| LibFn::Chi
| LibFn::FresnelS
| LibFn::FresnelC
| LibFn::AiryAi
| LibFn::AiryBi
| LibFn::AiryAiPrime
| LibFn::AiryBiPrime
| LibFn::EllipticK
| LibFn::EllipticE,
1,
) => {
let f = args[0];
let df = get_deriv(cache, f, arena);
if arena.is_zero_structural(df) {
return Some(arena.zero);
}
let outer = unary_rule(arena, f)?;
Some(arena.mul(&[outer, df]))
}
(LibFn::DirichletEta, 1) => {
let df = get_deriv(cache, args[0], arena);
if arena.is_zero_structural(df) {
return Some(arena.zero);
}
Some(formal(arena, args[0]))
}
(LibFn::ExpInt | LibFn::LowerGamma | LibFn::UpperGamma | LibFn::PolyLog, 2) => {
let (p, f) = (args[0], args[1]);
let dp = get_deriv(cache, p, arena);
let df = get_deriv(cache, f, arena);
if !arena.is_zero_structural(dp) {
let mut terms: SmallVec<[ExprId; 2]> = SmallVec::new();
terms.push(formal(arena, p));
if !arena.is_zero_structural(df) {
terms.push(formal(arena, f));
}
return Some(arena.add(&terms));
}
if arena.is_zero_structural(df) {
return Some(arena.zero);
}
let p_minus_1 = arena.sub(p, arena.one);
let outer = match name {
LibFn::ExpInt => {
let e = apply_named(arena, LibFn::ExpInt, &[p_minus_1, f]);
arena.neg(e)
}
LibFn::LowerGamma | LibFn::UpperGamma => {
let x_pow = arena.pow(f, p_minus_1);
let neg_f = arena.neg(f);
let e = arena.exp(neg_f);
let v = arena.mul(&[x_pow, e]);
if name == LibFn::LowerGamma {
v
} else {
arena.neg(v)
}
}
_ => {
let li = apply_named(arena, LibFn::PolyLog, &[p_minus_1, f]);
arena.div(li, f)
}
};
Some(arena.mul(&[outer, df]))
}
(LibFn::EllipticF, 2) => {
let (phi, m) = (args[0], args[1]);
let dphi = get_deriv(cache, phi, arena);
let dm = get_deriv(cache, m, arena);
let mut terms: SmallVec<[ExprId; 2]> = SmallVec::new();
if !arena.is_zero_structural(dphi) {
let s = arena.sin(phi);
let two = arena.int(2);
let s2 = arena.pow(s, two);
let ms2 = arena.mul(&[m, s2]);
let inner = arena.sub(arena.one, ms2);
let root = arena.sqrt(inner);
let outer = arena.div(arena.one, root);
terms.push(arena.mul(&[outer, dphi]));
}
if !arena.is_zero_structural(dm) {
terms.push(formal(arena, m));
}
Some(match terms.len() {
0 => arena.zero,
1 => terms[0],
_ => arena.add(&terms),
})
}
(LibFn::EllipticPi, 2) => {
let (n, m) = (args[0], args[1]);
let dn = get_deriv(cache, n, arena);
let dm = get_deriv(cache, m, arena);
if arena.is_zero_structural(dn) && arena.is_zero_structural(dm) {
return Some(arena.zero);
}
let k = apply_named(arena, LibFn::EllipticK, &[m]);
let e = apply_named(arena, LibFn::EllipticE, &[m]);
let pi_nm = apply_named(arena, LibFn::EllipticPi, &[n, m]);
let two = arena.int(2);
let m_minus_n = arena.sub(m, n);
let mut terms: SmallVec<[ExprId; 2]> = SmallVec::new();
if !arena.is_zero_structural(dn) {
let t1 = arena.mul(&[m_minus_n, k]);
let t1 = arena.div(t1, n);
let n2 = arena.pow(n, two);
let n2_minus_m = arena.sub(n2, m);
let t2 = arena.mul(&[n2_minus_m, pi_nm]);
let t2 = arena.div(t2, n);
let numer = arena.add(&[e, t1, t2]);
let n_minus_1 = arena.sub(n, arena.one);
let denom = arena.mul(&[two, m_minus_n, n_minus_1]);
let outer = arena.div(numer, denom);
terms.push(arena.mul(&[outer, dn]));
}
if !arena.is_zero_structural(dm) {
let m_minus_1 = arena.sub(m, arena.one);
let t = arena.div(e, m_minus_1);
let numer = arena.add(&[t, pi_nm]);
let n_minus_m = arena.sub(n, m);
let denom = arena.mul(&[two, n_minus_m]);
let outer = arena.div(numer, denom);
terms.push(arena.mul(&[outer, dm]));
}
Some(if terms.len() == 1 {
terms[0]
} else {
arena.add(&terms)
})
}
(LibFn::Gegenbauer | LibFn::AssocLegendre | LibFn::AssocLaguerre, 3)
| (LibFn::Jacobi, 4) => {
let x = args[args.len() - 1];
let params = &args[..args.len() - 1];
for &p in params {
let dp = get_deriv(cache, p, arena);
if !arena.is_zero_structural(dp) {
let mut terms: SmallVec<[ExprId; 4]> = SmallVec::new();
for &a in args {
let da = get_deriv(cache, a, arena);
if !arena.is_zero_structural(da) {
terms.push(formal(arena, a));
}
}
return Some(arena.add(&terms));
}
}
let dx = get_deriv(cache, x, arena);
if arena.is_zero_structural(dx) {
return Some(arena.zero);
}
let n = params[0];
let n_minus_1 = arena.sub(n, arena.one);
let outer = match name {
LibFn::Gegenbauer => {
let a = params[1];
let a_plus_1 = arena.add(&[a, arena.one]);
let c = apply_named(arena, LibFn::Gegenbauer, &[n_minus_1, a_plus_1, x]);
let two = arena.int(2);
arena.mul(&[two, a, c])
}
LibFn::Jacobi => {
let (a, b) = (params[1], params[2]);
let a_plus_1 = arena.add(&[a, arena.one]);
let b_plus_1 = arena.add(&[b, arena.one]);
let p = apply_named(arena, LibFn::Jacobi, &[n_minus_1, a_plus_1, b_plus_1, x]);
let s = arena.add(&[n, a, b, arena.one]);
let half = arena.rational(1, 2);
arena.mul(&[half, s, p])
}
LibFn::AssocLegendre => {
let m = params[1];
let pn = apply_named(arena, LibFn::AssocLegendre, &[n, m, x]);
let pn1 = apply_named(arena, LibFn::AssocLegendre, &[n_minus_1, m, x]);
let t1 = arena.mul(&[n, x, pn]);
let n_plus_m = arena.add(&[n, m]);
let t2 = arena.mul(&[n_plus_m, pn1]);
let numer = arena.sub(t1, t2);
let two = arena.int(2);
let x2 = arena.pow(x, two);
let denom = arena.sub(x2, arena.one);
arena.div(numer, denom)
}
_ => {
let a = params[1];
let a_plus_1 = arena.add(&[a, arena.one]);
let l = apply_named(arena, LibFn::AssocLaguerre, &[n_minus_1, a_plus_1, x]);
arena.neg(l)
}
};
Some(arena.mul(&[outer, dx]))
}
(LibFn::BetaInc | LibFn::BetaIncRegularized, 4) => {
let (a, b, x1, x2) = (args[0], args[1], args[2], args[3]);
let dx1 = get_deriv(cache, x1, arena);
let dx2 = get_deriv(cache, x2, arena);
let mut terms: SmallVec<[ExprId; 4]> = SmallVec::new();
for &p in &[a, b] {
let dp = get_deriv(cache, p, arena);
if !arena.is_zero_structural(dp) {
terms.push(formal(arena, p));
}
}
let integrand = |arena: &mut Arena, t: ExprId| -> ExprId {
let a_minus_1 = arena.sub(a, arena.one);
let b_minus_1 = arena.sub(b, arena.one);
let one_minus_t = arena.sub(arena.one, t);
let p1 = arena.pow(t, a_minus_1);
let p2 = arena.pow(one_minus_t, b_minus_1);
let v = arena.mul(&[p1, p2]);
if name == LibFn::BetaIncRegularized {
let beta = arena.beta(a, b);
arena.div(v, beta)
} else {
v
}
};
if !arena.is_zero_structural(dx2) {
let f = integrand(arena, x2);
terms.push(arena.mul(&[f, dx2]));
}
if !arena.is_zero_structural(dx1) {
let f = integrand(arena, x1);
let neg_f = arena.neg(f);
terms.push(arena.mul(&[neg_f, dx1]));
}
Some(match terms.len() {
0 => arena.zero,
1 => terms[0],
_ => arena.add(&terms),
})
}
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::base::arena::Arena;
fn sym(a: &mut Arena, name: &str) -> ExprId {
a.symbol(name)
}
fn display(a: &Arena, id: ExprId) -> String {
a.display(id).to_string()
}
#[test]
fn diff_constant_is_zero() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let five = a.int(5);
assert_eq!(diff(&mut a, five, x), a.zero);
}
#[test]
fn diff_pi_is_zero() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let pi = a.pi;
let result = diff(&mut a, pi, x);
assert_eq!(result, a.zero);
}
#[test]
fn diff_other_symbol_is_zero() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
assert_eq!(diff(&mut a, y, x), a.zero);
}
#[test]
fn diff_x_is_one() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
assert_eq!(diff(&mut a, x, x), a.one);
}
#[test]
fn diff_add() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let three = a.int(3);
let expr = a.add(&[x, three]);
let result = diff(&mut a, expr, x);
assert_eq!(result, a.one);
}
#[test]
fn diff_add_two_xs() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let expr = a.add(&[x, x]);
let result = diff(&mut a, expr, x);
let two = a.int(2);
assert_eq!(result, two);
}
#[test]
fn diff_mul_constant_times_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let three = a.int(3);
let expr = a.mul(&[three, x]);
let result = diff(&mut a, expr, x);
assert_eq!(display(&a, result), "3");
}
#[test]
fn diff_mul_x_times_y() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let expr = a.mul(&[x, y]);
let result = diff(&mut a, expr, x);
assert_eq!(result, y);
}
#[test]
fn diff_mul_x_times_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let expr = a.mul(&[x, x]);
let result = diff(&mut a, expr, x);
assert_eq!(display(&a, result), "2*x");
}
#[test]
fn diff_product_three_factors() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let z = sym(&mut a, "z");
let expr = a.mul(&[x, y, z]);
let result = diff(&mut a, expr, x);
assert_eq!(display(&a, result), "y*z");
}
#[test]
fn diff_x_squared() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let two = a.int(2);
let expr = a.pow(x, two);
let result = diff(&mut a, expr, x);
assert_eq!(display(&a, result), "2*x");
}
#[test]
fn diff_x_cubed() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let three = a.int(3);
let expr = a.pow(x, three);
let result = diff(&mut a, expr, x);
assert_eq!(display(&a, result), "3*x^2");
}
#[test]
fn diff_x_to_the_one() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let one = a.one;
let expr = a.pow(x, one);
assert_eq!(expr, x); let result = diff(&mut a, expr, x);
assert_eq!(result, a.one);
}
#[test]
fn diff_constant_power() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let two = a.int(2);
let three = a.int(3);
let expr = a.pow(two, three);
let result = diff(&mut a, expr, x);
assert_eq!(result, a.zero);
}
#[test]
fn diff_sin_x_squared() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let two = a.int(2);
let sin_x = a.sin(x);
let expr = a.pow(sin_x, two);
let result = diff(&mut a, expr, x);
let s = display(&a, result);
assert!(s.contains("2"), "should contain 2, got: {s}");
assert!(s.contains("sin"), "should contain sin, got: {s}");
assert!(s.contains("cos"), "should contain cos, got: {s}");
}
#[test]
fn diff_neg_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let expr = a.neg(x);
let result = diff(&mut a, expr, x);
assert_eq!(result, a.neg_one);
}
#[test]
fn diff_sin_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let expr = a.sin(x);
let result = diff(&mut a, expr, x);
assert_eq!(display(&a, result), "cos(x)");
}
#[test]
fn diff_cos_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let expr = a.cos(x);
let result = diff(&mut a, expr, x);
assert_eq!(display(&a, result), "-sin(x)");
}
#[test]
fn diff_tan_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let expr = a.tan(x);
let result = diff(&mut a, expr, x);
let s = display(&a, result);
assert!(s.contains("tan"), "should contain tan, got: {s}");
}
#[test]
fn diff_sin_chain_rule() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let two = a.int(2);
let x_sq = a.pow(x, two);
let expr = a.sin(x_sq);
let result = diff(&mut a, expr, x);
let s = display(&a, result);
assert!(s.contains("cos"), "should contain cos, got: {s}");
assert!(s.contains("2"), "should contain 2, got: {s}");
assert!(s.contains("x"), "should contain x, got: {s}");
}
#[test]
fn diff_exp_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let expr = a.exp(x);
let result = diff(&mut a, expr, x);
assert_eq!(display(&a, result), "exp(x)");
}
#[test]
fn diff_ln_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let expr = a.ln(x);
let result = diff(&mut a, expr, x);
assert_eq!(display(&a, result), "1/x");
}
#[test]
fn diff_exp_chain_rule() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let two = a.int(2);
let two_x = a.mul(&[two, x]);
let expr = a.exp(two_x);
let result = diff(&mut a, expr, x);
let s = display(&a, result);
assert!(s.contains("exp"), "should contain exp, got: {s}");
assert!(s.contains("2"), "should contain 2, got: {s}");
}
#[test]
fn diff_sqrt_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let expr = a.sqrt(x);
let result = diff(&mut a, expr, x);
let s = display(&a, result);
assert!(s.contains("sqrt") || s.contains("1/2"), "got: {s}");
}
#[test]
fn diff_polynomial() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let two = a.int(2);
let three = a.int(3);
let x3 = a.pow(x, three);
let x2 = a.pow(x, two);
let two_x2 = a.mul(&[two, x2]);
let five = a.int(5);
let expr = a.add(&[x3, two_x2, x, five]);
let result = diff(&mut a, expr, x);
let s = display(&a, result);
assert!(s.contains("3*x^2"), "should contain 3*x^2, got: {s}");
assert!(s.contains("4*x"), "should contain 4*x, got: {s}");
assert!(s.contains('1'), "should contain 1, got: {s}");
}
#[test]
fn diff_second_derivative() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let three = a.int(3);
let expr = a.pow(x, three);
let first = diff(&mut a, expr, x);
let second = diff(&mut a, first, x);
assert_eq!(display(&a, second), "6*x");
}
#[test]
fn diff_deep_expression_no_stack_overflow() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let mut expr = x;
for _ in 0..100 {
expr = a.sin(expr);
}
let result = diff(&mut a, expr, x);
assert_ne!(
result, a.zero,
"derivative of sin^100(x) should not be zero"
);
}
#[test]
fn diff_x_times_sin_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let sin_x = a.sin(x);
let expr = a.mul(&[x, sin_x]);
let result = diff(&mut a, expr, x);
let s = display(&a, result);
assert!(s.contains("sin"), "should contain sin, got: {s}");
assert!(s.contains("cos"), "should contain cos, got: {s}");
}
#[test]
fn diff_integral_fundamental_theorem() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let two = a.int(2);
let x2 = a.pow(x, two);
let integral = a.intern(crate::base::node::ExprNode::Integral(x2, x));
let result = diff(&mut a, integral, x);
assert_eq!(display(&a, result), "x^2");
}
#[test]
fn diff_integral_different_var_stays_unevaluated() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let integral = a.intern(crate::base::node::ExprNode::Integral(y, y));
let result = diff(&mut a, integral, x);
assert_eq!(display(&a, result), "Derivative(Integral(y, y), x)");
}
#[test]
fn diff_apply_constant_arg_is_zero() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let three = a.int(3);
let f_sid = a.symbols.intern("f");
let apply = a.intern(ExprNode::Apply(f_sid, smallvec::smallvec![three]));
let result = diff(&mut a, apply, x);
assert_eq!(result, a.zero, "d/dx(f(3)) should be 0");
}
#[test]
fn diff_apply_identity_is_formal_derivative() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let f_sid = a.symbols.intern("f");
let apply = a.intern(ExprNode::Apply(f_sid, smallvec::smallvec![x]));
let result = diff(&mut a, apply, x);
let s = display(&a, result);
assert!(
s.contains("Derivative") && s.contains("f(x)"),
"d/dx(f(x)) should be Derivative(f(x), x), got: {s}"
);
}
#[test]
fn diff_apply_chain_rule_x_squared() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let two = a.int(2);
let x_sq = a.pow(x, two);
let f_sid = a.symbols.intern("f");
let apply = a.intern(ExprNode::Apply(f_sid, smallvec::smallvec![x_sq]));
let result = diff(&mut a, apply, x);
let s = display(&a, result);
assert!(
s.contains('2') && s.contains('x') && s.contains("Derivative"),
"d/dx(f(x²)) should contain 2, x, Derivative, got: {s}"
);
}
#[test]
fn diff_apply_chain_sin_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let sin_x = a.sin(x);
let f_sid = a.symbols.intern("f");
let apply = a.intern(ExprNode::Apply(f_sid, smallvec::smallvec![sin_x]));
let result = diff(&mut a, apply, x);
let s = display(&a, result);
assert!(
s.contains("cos"),
"d/dx(f(sin(x))) should contain cos(x), got: {s}"
);
}
#[test]
fn diff_apply_other_symbol_is_zero() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let f_sid = a.symbols.intern("f");
let apply = a.intern(ExprNode::Apply(f_sid, smallvec::smallvec![y]));
let result = diff(&mut a, apply, x);
assert_eq!(result, a.zero, "d/dx(f(y)) should be 0");
}
#[test]
fn diff_known_apply_fibonacci_constant() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let ten = a.int(10);
let fib = a.fibonacci(ten);
let result = diff(&mut a, fib, x);
assert_eq!(result, a.zero, "d/dx(fibonacci(10)) should be 0");
}
#[test]
fn diff_known_apply_fibonacci_variable() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let fib = a.fibonacci(x);
let result = diff(&mut a, fib, x);
let s = display(&a, result);
assert!(
s.contains("Derivative") && s.contains("fibonacci"),
"d/dx(fibonacci(x)) should be formal Derivative, got: {s}"
);
}
#[test]
fn diff_apply_multiarg_chain_rule() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let two = a.int(2);
let x_sq = a.pow(x, two);
let f_sid = a.symbols.intern("f");
let apply = a.intern(ExprNode::Apply(f_sid, smallvec::smallvec![x, x_sq]));
let result = diff(&mut a, apply, x);
let s = display(&a, result);
assert!(
s.contains("Derivative"),
"d/dx(f(x, x²)) should contain Derivative terms, got: {s}"
);
}
#[test]
fn diff_apply_all_constant_args_is_zero() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let five = a.int(5);
let seven = a.int(7);
let g_sid = a.symbols.intern("g");
let apply = a.intern(ExprNode::Apply(g_sid, smallvec::smallvec![five, seven]));
let result = diff(&mut a, apply, x);
assert_eq!(result, a.zero, "d/dx(g(5,7)) should be 0");
}
#[test]
fn diff_apply_no_args_is_zero() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let h_sid = a.symbols.intern("h");
let apply = a.intern(ExprNode::Apply(h_sid, smallvec::smallvec![]));
let result = diff(&mut a, apply, x);
assert_eq!(result, a.zero, "d/dx(h()) should be 0");
}
#[test]
fn diff_with_deps_y_depends_on_x() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let mut deps = FxHashSet::default();
deps.insert(y);
let result = diff_with_deps(&mut a, y, x, &deps);
let s = display(&a, result);
assert!(
s.contains("Derivative") && s.contains('y') && s.contains('x'),
"d/dx(y) with deps={{y}} should be Derivative(y, x), got: {s}"
);
}
#[test]
fn diff_with_deps_y_squared() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let two = a.int(2);
let y_sq = a.pow(y, two);
let mut deps = FxHashSet::default();
deps.insert(y);
let result = diff_with_deps(&mut a, y_sq, x, &deps);
let s = display(&a, result);
assert!(
s.contains('2') && s.contains('y') && s.contains("Derivative"),
"d/dx(y²) with deps={{y}} should be 2*y*Derivative(y,x), got: {s}"
);
}
#[test]
fn diff_with_deps_implicit() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let two = a.int(2);
let x_sq = a.pow(x, two);
let y_sq = a.pow(y, two);
let expr = a.add(&[x_sq, y_sq]);
let mut deps = FxHashSet::default();
deps.insert(y);
let result = diff_with_deps(&mut a, expr, x, &deps);
let s = display(&a, result);
assert!(
s.contains('2') && s.contains('x') && s.contains("Derivative"),
"d/dx(x²+y²) with deps={{y}} should have 2*x and Derivative, got: {s}"
);
}
#[test]
fn diff_with_deps_sin_y() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let sin_y = a.sin(y);
let mut deps = FxHashSet::default();
deps.insert(y);
let result = diff_with_deps(&mut a, sin_y, x, &deps);
let s = display(&a, result);
assert!(
s.contains("cos") && s.contains("Derivative"),
"d/dx(sin(y)) with deps={{y}} should contain cos and Derivative, got: {s}"
);
}
#[test]
fn diff_with_deps_product_xy() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let xy = a.mul(&[x, y]);
let mut deps = FxHashSet::default();
deps.insert(y);
let result = diff_with_deps(&mut a, xy, x, &deps);
let s = display(&a, result);
assert!(
s.contains('y') && s.contains("Derivative"),
"d/dx(x*y) with deps={{y}} should contain y and Derivative, got: {s}"
);
}
#[test]
fn diff_with_deps_empty_is_zero() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let y = sym(&mut a, "y");
let deps = FxHashSet::default();
let result = diff_with_deps(&mut a, y, x, &deps);
assert_eq!(result, a.zero, "d/dx(y) with empty deps should be 0");
}
#[test]
fn diff_with_deps_matches_plain_diff() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let two = a.int(2);
let x_sq = a.pow(x, two);
let sin_x2 = a.sin(x_sq);
let deps = FxHashSet::default();
let r1 = diff_with_deps(&mut a, sin_x2, x, &deps);
let r2 = diff(&mut a, sin_x2, x);
assert_eq!(r1, r2, "diff_with_deps with empty deps should match diff");
}
#[test]
fn diff_with_deps_var_not_symbol_returns_zero() {
let mut a = Arena::new();
let three = a.int(3);
let y = sym(&mut a, "y");
let deps = FxHashSet::default();
let result = diff_with_deps(&mut a, y, three, &deps);
assert_eq!(result, a.zero);
}
fn real_sym(a: &mut Arena, name: &str) -> ExprId {
use crate::base::assumptions::{Assumptions, Props};
let id = a.symbol(name);
if let ExprNode::Symbol(sid) = a.node(id).clone() {
let mut asm = Assumptions::default();
asm.assert_true(Props::REAL);
a.set_symbol_assumptions(sid, asm);
}
id
}
#[test]
fn diff_named_constants_are_zero() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
for c in [a.euler_gamma, a.catalan, a.golden_ratio] {
assert_eq!(diff(&mut a, c, x), a.zero);
}
}
fn dd(a: &mut Arena, f: ExprId, x: ExprId) -> String {
let d = diff(a, f, x);
display(a, d)
}
#[test]
fn diff_special_functions() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let si = a.si(x);
assert_eq!(dd(&mut a, si, x), "sin(x)/x");
let ci = a.ci(x);
assert_eq!(dd(&mut a, ci, x), "cos(x)/x");
let ei = a.ei(x);
assert_eq!(dd(&mut a, ei, x), "exp(x)/x");
let li = a.li(x);
assert_eq!(dd(&mut a, li, x), "1/ln(x)");
let dg = a.digamma(x);
assert_eq!(dd(&mut a, dg, x), "polygamma(1, x)");
let two = a.int(2);
let pg = a.polygamma(two, x);
assert_eq!(dd(&mut a, pg, x), "polygamma(3, x)");
let z = a.zeta(x);
let dz = diff(&mut a, z, x);
assert!(matches!(a.node(dz), ExprNode::Derivative(_, _)));
let y = sym(&mut a, "y");
let kd = a.kronecker_delta(x, y);
assert_eq!(diff(&mut a, kd, x), a.zero);
let x2 = a.pow(x, two);
let si_x2 = a.si(x2);
assert_eq!(dd(&mut a, si_x2, x), "2*sin(x^2)/x");
}
#[test]
fn diff_complex_nodes_require_real_variable() {
let mut a = Arena::new();
let t = real_sym(&mut a, "t");
let z = sym(&mut a, "z");
let two = a.int(2);
let t2 = a.pow(t, two);
let f = a.mul(&[z, t2]);
let re_f = a.re(f);
assert_eq!(dd(&mut a, re_f, t), "2*t*re(z)");
let im_f = a.im(f);
assert_eq!(dd(&mut a, im_f, t), "2*t*im(z)");
let cf = a.conjugate(f);
assert_eq!(dd(&mut a, cf, t), "2*t*conjugate(z)");
let et = a.exp(t);
let g = a.mul(&[z, et]);
let ag = a.arg(g);
assert_eq!(diff(&mut a, ag, t), a.zero);
let x = sym(&mut a, "x");
let re_x = a.re(x);
let d = diff(&mut a, re_x, x);
assert!(matches!(a.node(d), ExprNode::Derivative(_, _)));
}
#[test]
fn diff_bessel_and_orthogonal_apply() {
let mut a = Arena::new();
let x = sym(&mut a, "x");
let nu = sym(&mut a, "nu");
let j = a.besselj(nu, x);
let dj = diff(&mut a, j, x);
let half = a.rational(1, 2);
let nm1 = a.sub(nu, a.one);
let np1 = a.add(&[nu, a.one]);
let jm = a.besselj(nm1, x);
let jp = a.besselj(np1, x);
let diffj = a.sub(jm, jp);
let expected = a.mul(&[half, diffj]);
assert_eq!(dj, expected);
let k = a.besselk(nu, x);
let dk = diff(&mut a, k, x);
let km = a.besselk(nm1, x);
let kp = a.besselk(np1, x);
let sumk = a.add(&[km, kp]);
let neg_half = a.rational(-1, 2);
let expected = a.mul(&[neg_half, sumk]);
assert_eq!(dk, expected);
let n = sym(&mut a, "n");
let t = a.chebyshev_t(n, x);
assert_eq!(dd(&mut a, t, x), "n*chebyshev_u(n - 1, x)");
let h = a.hermite(n, x);
assert_eq!(dd(&mut a, h, x), "2*n*hermite(n - 1, x)");
let jx = a.besselj(x, x);
let d = diff(&mut a, jx, x);
assert!(crate::base::walk::has_unevaluated(&a, d));
}
}