use crate::context::truncation_order;
use crate::da::Da;
use crate::error::{codes, dace_panic};
use crate::kernels::{multiply, weighted_sum};
pub(crate) fn evaluate_series(a: &Da, xf: &[f64]) -> Da {
let nocut = truncation_order();
assert!(
xf.len() > nocut as usize,
"series coefficient table too short"
);
let mut inon = a.clone();
inon.set_coefficient0(0, 0.0);
crate::context::set_truncation_order(1);
let mut inc = weighted_sum(&inon, xf[nocut as usize], &inon, 0.0);
inc = add_double(&inc, xf[(nocut - 1) as usize]);
let mut i = nocut as i64 - 2;
while i >= 0 {
crate::context::set_truncation_order(nocut - i as u32);
inc = multiply(&inon, &inc);
inc = add_double(&inc, xf[i as usize]);
i -= 1;
}
crate::context::set_truncation_order(nocut);
inc
}
pub(crate) fn add_double(a: &Da, c: f64) -> Da {
let mut r = a.clone();
let cons = r.cons();
r.set_coefficient0(0, cons + c);
r
}
pub(crate) fn double_subtract(a: &Da, c: f64) -> Da {
let mut r = weighted_sum(a, -1.0, a, 0.0);
let cons = r.cons();
r.set_coefficient0(0, cons + c);
r
}
pub(crate) fn subtract_double(a: &Da, c: f64) -> Da {
add_double(a, -c)
}
pub(crate) fn divide_double(a: &Da, c: f64) -> Da {
if c == 0.0 {
dace_panic(codes::DIVIDING_BY_ZERO, "Dividing by zero");
}
weighted_sum(a, 1.0 / c, a, 0.0)
}
pub(crate) fn minv(a: &Da) -> Da {
let a0 = a.cons();
if a0 == 0.0 {
dace_panic(codes::DIVIDING_BY_ZERO, "Dividing by zero");
}
let nocut = truncation_order();
if nocut < 5 {
minv0(a, a0)
} else {
crate::context::set_truncation_order(2);
let mut inc = minv0(a, a0);
let mut ord: u32 = 3;
while ord <= nocut {
crate::context::set_truncation_order(nocut.min(2 * ord - 1));
let temp = multiply(a, &inc);
let temp = double_subtract(&temp, 2.0);
inc = multiply(&inc, &temp);
ord *= 2;
}
crate::context::set_truncation_order(nocut);
inc
}
}
fn minv0(a: &Da, a0: f64) -> Da {
let scaled = weighted_sum(a, 1.0 / a0, a, 0.0);
let nocut = truncation_order();
let mut xf = vec![0.0; nocut as usize + 1];
xf[0] = 1.0 / a0;
for i in 1..xf.len() {
xf[i] = -xf[i - 1];
}
evaluate_series(&scaled, &xf)
}
pub fn trunc(a: &Da) -> Da {
let mut r = a.clone();
let c = r.cons().round_ties_even();
r.set_coefficient0(0, c);
r
}
pub fn round(a: &Da) -> Da {
let mut r = a.clone();
let c = r.cons().round();
r.set_coefficient0(0, c);
r
}
pub fn modulo(a: &Da, p: f64) -> Da {
let mut r = a.clone();
let c = r.cons() % p;
r.set_coefficient0(0, c);
r
}
pub fn powf(a: &Da, p: f64) -> Da {
if p == 0.0 {
return Da::constant(1.0);
}
if p.fract() == 0.0 && p.abs() <= i32::MAX as f64 {
return powi(a, p as i32);
}
let a0 = a.cons();
if a0 <= 0.0 {
dace_panic(
codes::NON_INTEGER_POWER_NON_POSITIVE,
"Non-integer power of non-positive DA",
);
}
let nocut = truncation_order();
let mut xf = vec![0.0; nocut as usize + 1];
xf[0] = a0.powf(p);
for i in 1..xf.len() {
xf[i] = xf[i - 1] / i as f64 * (p - (i - 1) as f64);
}
let scaled = divide_double(a, a0);
evaluate_series(&scaled, &xf)
}
pub fn powi(a: &Da, np: i32) -> Da {
match np {
0 => Da::constant(1.0),
1 => a.clone(),
-1 => minv(a),
_ => {
let abs_np = np.unsigned_abs();
let mut result = match abs_np {
2 => a.sqr(),
3 => multiply(a, &a.sqr()),
4 => a.sqr().sqr(),
_ => {
let mut itemp = a.clone();
let mut inc = Da::constant(1.0);
let mut inp = abs_np;
while inp > 0 {
if inp & 1 != 0 {
inc = multiply(&inc, &itemp);
}
inp >>= 1;
if inp > 0 {
itemp = itemp.sqr();
}
}
inc
}
};
if np < 0 {
result = minv(&result);
}
result
}
}
}
pub fn root(a: &Da, np: i32) -> Da {
if np == 0 {
dace_panic(codes::ZERO_TH_ROOT, "Zero-th root does not exist");
}
let a0 = a.cons();
let iodd = np.unsigned_abs() & 1;
if iodd == 0 && a0 <= 0.0 {
dace_panic(codes::EVEN_ROOT_NEGATIVE, "Even root of negative DA");
} else if iodd == 1 && a0 == 0.0 {
dace_panic(codes::ODD_ROOT_ZERO, "Odd root of zero DA");
}
let nocut = truncation_order();
let mut xf = vec![0.0; nocut as usize + 1];
let mut cr = 1.0 / f64::from(np);
xf[0] = a0.abs().powf(cr).copysign(a0);
for i in 1..xf.len() {
xf[i] = xf[i - 1] / i as f64 * cr;
cr -= 1.0;
}
let scaled = divide_double(a, a0);
evaluate_series(&scaled, &xf)
}
pub fn sqrt(a: &Da) -> Da {
root(a, 2)
}
pub fn isrt(a: &Da) -> Da {
root(a, -2)
}
pub fn cbrt(a: &Da) -> Da {
root(a, 3)
}
pub fn icrt(a: &Da) -> Da {
root(a, -3)
}
pub fn hypot(a: &Da, b: &Da) -> Da {
Da::assert_same_context(a, b);
root(&(a.sqr() + b.sqr()), 2)
}
pub fn exp(a: &Da) -> Da {
let nocut = truncation_order();
let mut xf = vec![0.0; nocut as usize + 1];
xf[0] = a.cons().exp();
for i in 1..xf.len() {
xf[i] = xf[i - 1] / i as f64;
}
evaluate_series(a, &xf)
}
pub fn log(a: &Da) -> Da {
let a0 = a.cons();
if a0 <= 0.0 {
dace_panic(
codes::LOG_NON_POSITIVE,
"Negative constant part in logarithm",
);
}
let nocut = truncation_order();
let mut xf = vec![0.0; nocut as usize + 1];
let scaled = divide_double(a, a0);
xf[0] = a0.ln();
xf[1] = 1.0;
for i in 2..xf.len() {
xf[i] = -xf[i - 1] / i as f64 * (i - 1) as f64;
}
evaluate_series(&scaled, &xf)
}
pub fn log_base(a: &Da, b: f64) -> Da {
if b <= 0.0 {
dace_panic(
codes::LOG_BASE_POSITIVE,
"Base of logarithm must be positive",
);
}
let l = log(a);
weighted_sum(&l, 1.0 / b.ln(), &l, 0.0)
}
pub fn log10(a: &Da) -> Da {
log_base(a, 10.0)
}
pub fn log2(a: &Da) -> Da {
log_base(a, 2.0)
}
pub fn sin(a: &Da) -> Da {
let nocut = truncation_order();
let a0 = a.cons();
let mut xf = vec![0.0; nocut as usize + 1];
xf[0] = a0.sin();
xf[1] = a0.cos();
for i in 2..xf.len() {
xf[i] = -xf[i - 2] / (i * (i - 1)) as f64;
}
evaluate_series(a, &xf)
}
pub fn cos(a: &Da) -> Da {
let nocut = truncation_order();
let a0 = a.cons();
let mut xf = vec![0.0; nocut as usize + 1];
xf[0] = a0.cos();
xf[1] = -a0.sin();
for i in 2..xf.len() {
xf[i] = -xf[i - 2] / (i * (i - 1)) as f64;
}
evaluate_series(a, &xf)
}
pub fn tan(a: &Da) -> Da {
if a.cons().cos() == 0.0 {
dace_panic(codes::COS_ZERO_IN_TANGENT, "Cosine is zero in tangent");
}
let s = sin(a);
let c = cos(a);
divide_da(&s, &c)
}
pub fn asin(a: &Da) -> Da {
if a.cons().abs() >= 1.0 {
dace_panic(codes::OUT_OF_DOMAIN, "Out of domain");
}
let d = double_subtract(&a.sqr(), 1.0);
let d = sqrt(&d);
let q = divide_da(a, &d);
atan(&q)
}
pub fn acos(a: &Da) -> Da {
if a.cons().abs() >= 1.0 {
dace_panic(codes::OUT_OF_DOMAIN, "Out of domain");
}
let s = asin(a);
double_subtract(&s, std::f64::consts::FRAC_PI_2)
}
pub fn atan(a: &Da) -> Da {
let nocut = truncation_order();
let a0 = a.cons();
let mut xf = vec![0.0; nocut as usize + 1];
let iarg = {
let denom = add_double(&weighted_sum(a, a0, a, 0.0), 1.0);
let num = subtract_double(a, a0);
divide_da(&num, &denom)
};
let mut s = 1.0;
xf[0] = a0.atan();
let mut i = 1;
while i < xf.len() {
xf[i] = s / i as f64;
s = -s;
i += 2;
}
evaluate_series(&iarg, &xf)
}
pub fn atan2(y: &Da, x: &Da) -> Da {
Da::assert_same_context(y, x);
let cx = x.cons();
let cy = y.cons();
if cx == 0.0 && cy == 0.0 {
return Da::constant(0.0);
}
if cy.abs() > cx.abs() {
let t = atan(÷_da(x, y));
if cy < 0.0 {
double_subtract(&t, -std::f64::consts::FRAC_PI_2)
} else {
double_subtract(&t, std::f64::consts::FRAC_PI_2)
}
} else {
let t = atan(÷_da(y, x));
if cx < 0.0 {
if cy > 0.0 {
add_double(&t, std::f64::consts::PI)
} else {
add_double(&t, -std::f64::consts::PI)
}
} else {
t
}
}
}
pub(crate) fn divide_da(a: &Da, b: &Da) -> Da {
multiply(a, &minv(b))
}
pub fn sinh(a: &Da) -> Da {
let nocut = truncation_order();
let a0 = a.cons();
let mut xf = vec![0.0; nocut as usize + 1];
xf[0] = a0.sinh();
xf[1] = a0.cosh();
for i in 2..xf.len() {
xf[i] = xf[i - 2] / (i * (i - 1)) as f64;
}
evaluate_series(a, &xf)
}
pub fn cosh(a: &Da) -> Da {
let nocut = truncation_order();
let a0 = a.cons();
let mut xf = vec![0.0; nocut as usize + 1];
xf[0] = a0.cosh();
xf[1] = a0.sinh();
for i in 2..xf.len() {
xf[i] = xf[i - 2] / (i * (i - 1)) as f64;
}
evaluate_series(a, &xf)
}
pub fn tanh(a: &Da) -> Da {
let a0 = a.cons();
if a0 > 0.0 {
let t = exp(&weighted_sum(a, -2.0, a, 0.0));
let denom = add_double(&t, 1.0);
let num = double_subtract(&t, 1.0);
divide_da(&num, &denom)
} else {
let t = exp(&weighted_sum(a, 2.0, a, 0.0));
let denom = add_double(&t, 1.0);
let num = add_double(&t, -1.0);
divide_da(&num, &denom)
}
}
pub fn asinh(a: &Da) -> Da {
let s = sqrt(&add_double(&a.sqr(), 1.0));
log(&(a.clone() + s))
}
pub fn acosh(a: &Da) -> Da {
if a.cons() <= 1.0 {
dace_panic(codes::OUT_OF_DOMAIN, "Out of domain");
}
let s = sqrt(&subtract_double(&a.sqr(), 1.0));
log(&(a.clone() + s))
}
pub fn atanh(a: &Da) -> Da {
if a.cons().abs() >= 1.0 {
dace_panic(codes::OUT_OF_DOMAIN, "Out of domain");
}
let num = add_double(a, 1.0);
let den = double_subtract(a, 1.0);
let q = log(÷_da(&num, &den));
weighted_sum(&q, 0.5, &q, 0.0)
}
pub fn erf(a: &Da) -> Da {
let nocut = truncation_order();
let a0 = a.cons();
let mut xf = vec![0.0; nocut as usize + 1];
let mut factor = 2.0 * (-a0 * a0).exp() / std::f64::consts::PI.sqrt();
xf[0] = puruspe::erf(a0);
xf[1] = factor;
let mut hi2 = 1.0; let mut hi1 = 2.0 * a0; for (i, item) in xf.iter_mut().enumerate().skip(2) {
factor /= -(i as f64);
*item = factor * hi1;
let temp = 2.0 * a0 * hi1 - 2.0 * (i - 1) as f64 * hi2;
hi2 = hi1;
hi1 = temp;
}
evaluate_series(a, &xf)
}
pub fn erfc(a: &Da) -> Da {
let nocut = truncation_order();
let a0 = a.cons();
let mut xf = vec![0.0; nocut as usize + 1];
let mut factor = -2.0 * (-a0 * a0).exp() / std::f64::consts::PI.sqrt();
xf[0] = puruspe::erfc(a0);
xf[1] = factor;
let mut hi2 = 1.0;
let mut hi1 = 2.0 * a0;
for (i, item) in xf.iter_mut().enumerate().skip(2) {
factor /= -(i as f64);
*item = factor * hi1;
let temp = 2.0 * a0 * hi1 - 2.0 * (i - 1) as f64 * hi2;
hi2 = hi1;
hi1 = temp;
}
evaluate_series(a, &xf)
}
macro_rules! method_form {
($($name:ident),* $(,)?) => {
$(
#[doc = concat!("Method form of [`", stringify!($name), "`](crate::elementary::", stringify!($name), ").")]
pub fn $name(&self) -> Da {
$name(self)
}
)*
};
}
impl Da {
method_form!(
exp, log, log10, log2, sin, cos, tan, asin, acos, atan, sinh, cosh, tanh, asinh, acosh,
atanh, erf, erfc, sqrt, isrt, cbrt, icrt,
);
pub fn powf(&self, p: f64) -> Da {
powf(self, p)
}
pub fn powi(&self, np: i32) -> Da {
powi(self, np)
}
pub fn root(&self, np: i32) -> Da {
root(self, np)
}
pub fn log_base(&self, b: f64) -> Da {
log_base(self, b)
}
pub fn hypot(&self, b: &Da) -> Da {
hypot(self, b)
}
pub fn atan2(&self, x: &Da) -> Da {
atan2(self, x)
}
pub fn modulo(&self, p: f64) -> Da {
modulo(self, p)
}
pub fn trunc(&self) -> Da {
trunc(self)
}
pub fn round(&self) -> Da {
round(self)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::test_support::CONTEXT_LOCK;
#[test]
fn taylor_coefficients() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(6, 1).unwrap();
let x = Da::variable(1);
let s = sin(&x);
assert!((s.get_coefficient(&[1]) - 1.0).abs() < 1e-15);
assert!((s.get_coefficient(&[3]) + 1.0 / 6.0).abs() < 1e-15);
assert!((s.get_coefficient(&[5]) - 1.0 / 120.0).abs() < 1e-15);
assert_eq!(s.size(), 3);
let e = exp(&x);
let mut fact = 1.0;
for k in 0..=6u32 {
if k > 0 {
fact *= f64::from(k);
}
assert!(
(e.get_coefficient(&[k]) - 1.0 / fact).abs() < 1e-15,
"exp coeff {k}"
);
}
let a = atan(&x);
assert!((a.get_coefficient(&[1]) - 1.0).abs() < 1e-15);
assert!((a.get_coefficient(&[3]) + 1.0 / 3.0).abs() < 1e-15);
assert!((a.get_coefficient(&[5]) - 1.0 / 5.0).abs() < 1e-15);
let u = 1.0 + x.clone();
let r = sqrt(&(u.clone() * u.clone()));
assert!((r.cons() - 1.0).abs() < 1e-12);
assert!((r.get_coefficient(&[1]) - 1.0).abs() < 1e-12);
let p3 = x.clone().powi(3);
assert!((p3.get_coefficient(&[3]) - 1.0).abs() < 1e-15);
assert_eq!(p3.size(), 1);
}
#[test]
fn identities_two_vars() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(8, 2).unwrap();
let x = Da::variable(1);
let y = Da::variable(2);
let f = 0.7 + 1.3 * x.clone() - 0.4 * y.clone();
let s2 = sin(&f).sqr() + cos(&f).sqr();
assert!((s2.cons() - 1.0).abs() < 1e-12);
for m in s2.iter_monomials() {
if m.order() == 0 {
continue;
}
assert!(m.c.abs() < 1e-12, "sin^2+cos^2 residual at {:?}", m.jj);
}
let lf = log(&exp(&f));
for m in lf.iter_monomials() {
let expect = f.get_coefficient(&m.jj);
assert!(
(m.c - expect).abs() <= 1e-12 * expect.abs().max(1.0),
"log(exp) at {:?}: {} vs {}",
m.jj,
m.c,
expect
);
}
let th = tanh(&f);
let sh = sinh(&f);
let ch = cosh(&f);
let q = divide_da(&sh, &ch);
for (m1, m2) in th.iter_monomials().zip(q.iter_monomials()) {
assert_eq!(m1.jj, m2.jj);
assert!((m1.c - m2.c).abs() <= 1e-11 * m1.c.abs().max(1.0));
}
let g = 0.3 * x.clone();
let asg = asin(&g);
let d = asg.deriv(1);
let expect = weighted_sum(&isrt(&(1.0 - g.clone() * g.clone())), 0.3, &Da::new(), 0.0);
for (m1, m2) in d.iter_monomials().zip(expect.iter_monomials()) {
assert_eq!(m1.jj, m2.jj);
assert!((m1.c - m2.c).abs() <= 1e-11 * m1.c.abs().max(1.0));
}
let ang = 0.5 + 0.2 * x.clone();
let r = 1.2 + 0.3 * y.clone();
let yv = r.clone() * sin(&ang.clone());
let xv = r.clone() * cos(&ang.clone());
let a2 = atan2(&yv, &xv);
for m in a2.iter_monomials() {
let expect = ang.get_coefficient(&m.jj);
assert!(
(m.c - expect).abs() <= 1e-11 * expect.abs().max(1.0),
"atan2 at {:?}: {} vs {}",
m.jj,
m.c,
expect
);
}
let z = 0.0 + x.clone();
let ez = erf(&z);
assert_eq!(ez.cons(), 0.0);
assert!((ez.get_coefficient(&[1, 0]) - 2.0 / std::f64::consts::PI.sqrt()).abs() < 1e-14);
let ec = erfc(&z);
assert!((ec.cons() - 1.0).abs() < 1e-14);
assert!((ec.get_coefficient(&[1, 0]) + 2.0 / std::f64::consts::PI.sqrt()).abs() < 1e-14);
let w = 2.5 + x.clone();
assert!((w.trunc().cons() - 2.0).abs() < 1e-15);
assert!((w.round().cons() - 3.0).abs() < 1e-15);
assert!((w.modulo(2.0).cons() - 0.5).abs() < 1e-12);
assert!((log10(&(Da::constant(100.0))).cons() - 2.0).abs() < 1e-12);
assert!((log2(&(Da::constant(8.0))).cons() - 3.0).abs() < 1e-12);
assert!((log_base(&(Da::constant(8.0)), 2.0).cons() - 3.0).abs() < 1e-12);
assert!((cbrt(&(Da::constant(27.0))).cons() - 3.0).abs() < 1e-12);
assert!((isrt(&(Da::constant(4.0))).cons() - 0.5).abs() < 1e-12);
assert!((icrt(&(Da::constant(8.0))).cons() - 0.5).abs() < 1e-12);
assert!((hypot(&(Da::constant(3.0)), &Da::constant(4.0)).cons() - 5.0).abs() < 1e-12);
assert!((powf(&(Da::constant(2.0)), 10.0).cons() - 1024.0).abs() < 1e-9);
assert!((acosh(&(Da::constant(2.0))).cons() - 2.0f64.acosh()).abs() < 1e-12);
assert!((asinh(&(Da::constant(1.5))).cons() - 1.5f64.asinh()).abs() < 1e-12);
assert!((atanh(&(Da::constant(0.5))).cons() - 0.5f64.atanh()).abs() < 1e-12);
}
#[test]
fn domain_panics() {
let _g = CONTEXT_LOCK.lock();
crate::context::init(6, 2).unwrap();
fn expect_code(f: impl FnOnce() -> Da + std::panic::UnwindSafe, code: u32) {
let err = std::panic::catch_unwind(f).expect_err("must panic");
let e = err
.downcast_ref::<crate::error::DaceError>()
.expect("DaceError payload");
assert_eq!(e.code, code, "{}", e);
}
expect_code(|| log(&Da::constant(-1.0)), codes::LOG_NON_POSITIVE);
expect_code(|| log(&Da::constant(0.0)), codes::LOG_NON_POSITIVE);
expect_code(|| Da::constant(0.0).minv(), codes::DIVIDING_BY_ZERO);
expect_code(|| sqrt(&Da::constant(-4.0)), codes::EVEN_ROOT_NEGATIVE);
expect_code(|| root(&Da::constant(1.0), 0), codes::ZERO_TH_ROOT);
expect_code(|| root(&Da::constant(0.0), 3), codes::ODD_ROOT_ZERO);
expect_code(
|| powf(&Da::constant(-2.0), 0.5),
codes::NON_INTEGER_POWER_NON_POSITIVE,
);
expect_code(|| asin(&Da::constant(1.0)), codes::OUT_OF_DOMAIN);
expect_code(|| acos(&Da::constant(-1.5)), codes::OUT_OF_DOMAIN);
expect_code(|| acosh(&Da::constant(1.0)), codes::OUT_OF_DOMAIN);
expect_code(|| atanh(&Da::constant(-1.0)), codes::OUT_OF_DOMAIN);
expect_code(
|| log_base(&Da::constant(2.0), -1.0),
codes::LOG_BASE_POSITIVE,
);
expect_code(
|| Da::constant(1.0) / Da::variable(1),
codes::DIVIDING_BY_ZERO,
);
}
}