use crate::lower::LoweredOp;
use std::sync::Arc;
const INDET_EPS: f64 = 1e-6;
const BIG: f64 = 1e12;
const LHOPITAL_MAX: usize = 8;
const CAUCHY_TOL: f64 = 1e-7;
const CAUCHY_FACTOR: f64 = 1000.0;
const H_LADDER: [f64; 5] = [1e-2, 1e-4, 1e-6, 1e-8, 1e-10];
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum LimitPoint {
Finite(f64),
PosInf,
NegInf,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum LimitResult {
Finite(f64),
PosInf,
NegInf,
DoesNotExist,
Indeterminate,
}
fn eval_at_wrt(op: &LoweredOp, wrt: usize, x: f64) -> f64 {
let needed = (wrt + 1).max(op.count_vars()).max(1);
let mut vars = vec![0.0_f64; needed];
vars[wrt] = x;
op.eval(&vars)
}
fn substitute(op: &LoweredOp, wrt: usize, replacement: &LoweredOp) -> LoweredOp {
match op {
LoweredOp::Var(i) => {
if *i == wrt {
replacement.clone()
} else {
op.clone()
}
}
LoweredOp::Const(_) | LoweredOp::NamedConst(_) => op.clone(),
LoweredOp::Neg(x) => LoweredOp::Neg(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Exp(x) => LoweredOp::Exp(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Ln(x) => LoweredOp::Ln(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Sin(x) => LoweredOp::Sin(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Cos(x) => LoweredOp::Cos(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Tan(x) => LoweredOp::Tan(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Sinh(x) => LoweredOp::Sinh(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Cosh(x) => LoweredOp::Cosh(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Tanh(x) => LoweredOp::Tanh(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Arcsin(x) => LoweredOp::Arcsin(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Arccos(x) => LoweredOp::Arccos(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Arctan(x) => LoweredOp::Arctan(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Arcsinh(x) => LoweredOp::Arcsinh(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Arccosh(x) => LoweredOp::Arccosh(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Arctanh(x) => LoweredOp::Arctanh(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Erf(x) => LoweredOp::Erf(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::LGamma(x) => LoweredOp::LGamma(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Digamma(x) => LoweredOp::Digamma(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Trigamma(x) => LoweredOp::Trigamma(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Ei(x) => LoweredOp::Ei(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Si(x) => LoweredOp::Si(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Ci(x) => LoweredOp::Ci(Arc::new(substitute(x, wrt, replacement))),
LoweredOp::Add(a, b) => LoweredOp::Add(
Arc::new(substitute(a, wrt, replacement)),
Arc::new(substitute(b, wrt, replacement)),
),
LoweredOp::Sub(a, b) => LoweredOp::Sub(
Arc::new(substitute(a, wrt, replacement)),
Arc::new(substitute(b, wrt, replacement)),
),
LoweredOp::Mul(a, b) => LoweredOp::Mul(
Arc::new(substitute(a, wrt, replacement)),
Arc::new(substitute(b, wrt, replacement)),
),
LoweredOp::Div(a, b) => LoweredOp::Div(
Arc::new(substitute(a, wrt, replacement)),
Arc::new(substitute(b, wrt, replacement)),
),
LoweredOp::Pow(a, b) => LoweredOp::Pow(
Arc::new(substitute(a, wrt, replacement)),
Arc::new(substitute(b, wrt, replacement)),
),
}
}
fn probe_side(op: &LoweredOp, wrt: usize, c: f64, from_right: bool) -> f64 {
let vals: Vec<f64> = H_LADDER
.iter()
.map(|&h| {
let x = if from_right { c + h } else { c - h };
eval_at_wrt(op, wrt, x)
})
.collect();
let last = match vals.last() {
Some(&v) => v,
None => return f64::NAN,
};
let n = vals.len();
let second_last = if n >= 2 { vals[n - 2] } else { f64::NAN };
if last.is_infinite() && second_last.is_infinite() {
return if last.signum() == second_last.signum() {
last
} else {
f64::NAN };
}
if last.is_infinite()
&& second_last.is_finite()
&& second_last.abs() > 1.0
&& last.signum() == second_last.signum()
{
return last;
}
if last.is_nan() || second_last.is_nan() {
return f64::NAN;
}
if last.is_finite() && second_last.is_finite() {
if last.abs() > second_last.abs() * 10.0
&& last.abs() > 1.0
&& last.signum() == second_last.signum()
{
let first = vals.first().copied().unwrap_or(f64::NAN);
if first.is_finite()
&& first.abs() > 0.01
&& last.abs() > first.abs() * 1000.0
&& last.signum() == first.signum()
{
return if last > 0.0 {
f64::INFINITY
} else {
f64::NEG_INFINITY
};
}
}
let diff = (last - second_last).abs();
let scale = last.abs().max(second_last.abs()).max(1.0);
if diff <= CAUCHY_TOL * CAUCHY_FACTOR * scale {
return last;
}
}
f64::NAN
}
fn classify_probes(right: f64, left: f64) -> LimitResult {
if right.is_nan() || left.is_nan() {
return LimitResult::DoesNotExist;
}
match (right.is_infinite(), left.is_infinite()) {
(true, true) => {
if right > 0.0 && left > 0.0 {
LimitResult::PosInf
} else if right < 0.0 && left < 0.0 {
LimitResult::NegInf
} else {
LimitResult::DoesNotExist }
}
(false, false) => {
let scale = right.abs().max(left.abs()).max(1.0);
if (right - left).abs() <= CAUCHY_TOL * CAUCHY_FACTOR * scale {
LimitResult::Finite((right + left) / 2.0)
} else {
LimitResult::DoesNotExist
}
}
_ => LimitResult::DoesNotExist, }
}
fn limit_at_finite(op: &LoweredOp, wrt: usize, c: f64, lhopital_count: usize) -> LimitResult {
let v = eval_at_wrt(op, wrt, c);
if v.is_finite() {
let right = probe_side(op, wrt, c, true);
let left = probe_side(op, wrt, c, false);
if right.is_nan() || left.is_nan() {
return LimitResult::DoesNotExist;
}
if right.is_infinite() || left.is_infinite() {
return match (right.is_infinite(), left.is_infinite()) {
(true, true) if right.signum() == left.signum() => {
if right > 0.0 {
LimitResult::PosInf
} else {
LimitResult::NegInf
}
}
_ => LimitResult::DoesNotExist,
};
}
let scale = right.abs().max(left.abs()).max(v.abs()).max(1.0);
let tol = CAUCHY_TOL * CAUCHY_FACTOR * scale;
if (right - v).abs() <= tol && (left - v).abs() <= tol {
return LimitResult::Finite(v);
}
if (right - left).abs() <= tol {
return LimitResult::Finite((right + left) / 2.0);
}
return LimitResult::DoesNotExist;
}
if lhopital_count < LHOPITAL_MAX {
if let LoweredOp::Div(num, den) = op {
let num_val = eval_at_wrt(num, wrt, c);
let den_val = eval_at_wrt(den, wrt, c);
let zero_over_zero = num_val.is_finite()
&& num_val.abs() < INDET_EPS
&& den_val.is_finite()
&& den_val.abs() < INDET_EPS;
let inf_over_inf = num_val.abs() > BIG && den_val.abs() > BIG;
if zero_over_zero || inf_over_inf {
let new_num = num.grad(wrt);
let new_den = den.grad(wrt);
let new_expr = LoweredOp::Div(Arc::new(new_num), Arc::new(new_den)).simplify();
return limit_inner(&new_expr, wrt, LimitPoint::Finite(c), lhopital_count + 1);
}
}
}
let right = probe_side(op, wrt, c, true);
let left = probe_side(op, wrt, c, false);
classify_probes(right, left)
}
fn limit_at_finite_one_sided(
op: &LoweredOp,
wrt: usize,
c: f64,
from_right: bool,
lhopital_count: usize,
) -> LimitResult {
let v = eval_at_wrt(op, wrt, c);
if v.is_finite() {
let side = probe_side(op, wrt, c, from_right);
if side.is_nan() {
return LimitResult::DoesNotExist;
}
if side.is_infinite() {
return if side > 0.0 {
LimitResult::PosInf
} else {
LimitResult::NegInf
};
}
let scale = side.abs().max(v.abs()).max(1.0);
if (side - v).abs() <= CAUCHY_TOL * CAUCHY_FACTOR * scale {
return LimitResult::Finite(v);
}
return LimitResult::Finite(side);
}
if lhopital_count < LHOPITAL_MAX {
if let LoweredOp::Div(num, den) = op {
let num_val = eval_at_wrt(num, wrt, c);
let den_val = eval_at_wrt(den, wrt, c);
let zero_over_zero = num_val.is_finite()
&& num_val.abs() < INDET_EPS
&& den_val.is_finite()
&& den_val.abs() < INDET_EPS;
let inf_over_inf = num_val.abs() > BIG && den_val.abs() > BIG;
if zero_over_zero || inf_over_inf {
let new_num = num.grad(wrt);
let new_den = den.grad(wrt);
let new_expr = LoweredOp::Div(Arc::new(new_num), Arc::new(new_den)).simplify();
return limit_at_finite_one_sided(
&new_expr,
wrt,
c,
from_right,
lhopital_count + 1,
);
}
}
}
let side = probe_side(op, wrt, c, from_right);
if side.is_nan() {
LimitResult::Indeterminate
} else if side.is_infinite() {
if side > 0.0 {
LimitResult::PosInf
} else {
LimitResult::NegInf
}
} else {
LimitResult::Finite(side)
}
}
fn limit_inner(
op: &LoweredOp,
wrt: usize,
point: LimitPoint,
lhopital_count: usize,
) -> LimitResult {
match point {
LimitPoint::Finite(c) => limit_at_finite(op, wrt, c, lhopital_count),
LimitPoint::PosInf => {
let inv_t = LoweredOp::Div(
Arc::new(LoweredOp::Const(1.0)),
Arc::new(LoweredOp::Var(wrt)),
);
let substituted = substitute(op, wrt, &inv_t).simplify();
limit_at_finite_one_sided(&substituted, wrt, 0.0, true, 0)
}
LimitPoint::NegInf => {
let neg_inv_t = LoweredOp::Div(
Arc::new(LoweredOp::Neg(Arc::new(LoweredOp::Const(1.0)))),
Arc::new(LoweredOp::Var(wrt)),
);
let substituted = substitute(op, wrt, &neg_inv_t).simplify();
limit_at_finite_one_sided(&substituted, wrt, 0.0, true, 0)
}
}
}
impl LoweredOp {
pub fn limit(&self, wrt: usize, point: LimitPoint) -> LimitResult {
limit_inner(&self.simplify(), wrt, point, 0)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn var() -> LoweredOp {
LoweredOp::Var(0)
}
fn c(v: f64) -> LoweredOp {
LoweredOp::Const(v)
}
fn div(a: LoweredOp, b: LoweredOp) -> LoweredOp {
LoweredOp::Div(Arc::new(a), Arc::new(b))
}
fn sub(a: LoweredOp, b: LoweredOp) -> LoweredOp {
LoweredOp::Sub(Arc::new(a), Arc::new(b))
}
fn add(a: LoweredOp, b: LoweredOp) -> LoweredOp {
LoweredOp::Add(Arc::new(a), Arc::new(b))
}
fn sin(a: LoweredOp) -> LoweredOp {
LoweredOp::Sin(Arc::new(a))
}
fn cos(a: LoweredOp) -> LoweredOp {
LoweredOp::Cos(Arc::new(a))
}
fn exp(a: LoweredOp) -> LoweredOp {
LoweredOp::Exp(Arc::new(a))
}
fn pow(a: LoweredOp, b: LoweredOp) -> LoweredOp {
LoweredOp::Pow(Arc::new(a), Arc::new(b))
}
fn assert_finite(result: LimitResult, expected: f64, tol: f64) {
match result {
LimitResult::Finite(v) => {
assert!(
(v - expected).abs() <= tol,
"expected Finite({expected}), got Finite({v}); |diff|={}",
(v - expected).abs()
);
}
other => panic!("expected Finite({expected}), got {other:?}"),
}
}
#[test]
fn test_sinc_at_zero() {
assert_finite(
div(sin(var()), var()).limit(0, LimitPoint::Finite(0.0)),
1.0,
0.001,
);
}
#[test]
fn test_one_minus_cos_over_x_sq() {
let expr = div(sub(c(1.0), cos(var())), pow(var(), c(2.0)));
assert_finite(expr.limit(0, LimitPoint::Finite(0.0)), 0.5, 0.001);
}
#[test]
fn test_exp_minus_one_over_x() {
let expr = div(sub(exp(var()), c(1.0)), var());
assert_finite(expr.limit(0, LimitPoint::Finite(0.0)), 1.0, 0.001);
}
#[test]
fn test_x_sq_minus_1_over_x_minus_1() {
let expr = div(sub(pow(var(), c(2.0)), c(1.0)), sub(var(), c(1.0)));
assert_finite(expr.limit(0, LimitPoint::Finite(1.0)), 2.0, 0.001);
}
#[test]
fn test_one_over_x_at_zero_dne() {
let expr = div(c(1.0), var());
assert_eq!(
expr.limit(0, LimitPoint::Finite(0.0)),
LimitResult::DoesNotExist
);
}
#[test]
fn test_sin_one_over_x_oscillates() {
let expr = sin(div(c(1.0), var()));
assert_eq!(
expr.limit(0, LimitPoint::Finite(0.0)),
LimitResult::DoesNotExist
);
}
#[test]
fn test_one_over_x_at_pos_inf() {
assert_finite(div(c(1.0), var()).limit(0, LimitPoint::PosInf), 0.0, 0.01);
}
#[test]
fn test_x_over_x_plus_1_at_inf() {
let expr = div(var(), add(var(), c(1.0)));
assert_finite(expr.limit(0, LimitPoint::PosInf), 1.0, 0.01);
}
#[test]
fn test_one_plus_inv_x_pow_x() {
let expr = pow(add(c(1.0), div(c(1.0), var())), var());
let result = expr.limit(0, LimitPoint::PosInf);
match result {
LimitResult::Finite(v) => {
assert!(
(v - std::f64::consts::E).abs() < 0.05,
"(1+1/x)^x at +∞: expected ≈e, got {v}"
);
}
other => panic!("(1+1/x)^x at +∞: expected Finite(≈e), got {other:?}"),
}
}
#[test]
fn test_exp_at_neg_inf() {
assert_finite(exp(var()).limit(0, LimitPoint::NegInf), 0.0, 0.01);
}
#[test]
fn test_limit_terminates() {
use std::time::Instant;
let expr = sin(div(c(1.0), var()));
let start = Instant::now();
let result = expr.limit(0, LimitPoint::Finite(0.0));
let elapsed = start.elapsed();
assert!(
elapsed.as_secs() < 10,
"limit computation took too long: {elapsed:?}"
);
assert_eq!(result, LimitResult::DoesNotExist);
}
}