use astro_float::{BigFloat, Consts, RoundingMode};
use num_bigint::BigInt;
use num_rational::Ratio;
use num_traits::{One, Signed, Zero};
use super::dense::Poly;
use crate::transforms::evalf::{c_add, c_div, c_from_real, c_mul, c_one, c_sub, c_zero};
type Complex = (BigFloat, BigFloat);
fn poly_eval_complex(
coeffs: &[Ratio<BigInt>],
z: &Complex,
prec: usize,
rm: RoundingMode,
) -> Complex {
if coeffs.is_empty() {
return c_zero(prec);
}
let mut result = c_zero(prec);
for c in coeffs.iter().rev() {
result = c_mul(&result, z, prec, rm);
let c_re = ratio_to_bigfloat(c, prec);
let c_complex = c_from_real(c_re, prec);
result = c_add(&result, &c_complex, prec, rm);
}
result
}
fn ratio_to_bigfloat(r: &Ratio<BigInt>, prec: usize) -> BigFloat {
let numer_f = bigint_to_bigfloat(r.numer(), prec);
let denom_f = bigint_to_bigfloat(r.denom(), prec);
if denom_f.is_zero() {
return BigFloat::new(prec);
}
numer_f.div(&denom_f, prec, RoundingMode::None)
}
fn bigint_to_bigfloat(n: &BigInt, prec: usize) -> BigFloat {
if let Ok(small) = i64::try_from(n) {
return BigFloat::from_i64(small, prec);
}
if let Ok(medium) = i128::try_from(n) {
return BigFloat::from_i128(medium, prec);
}
let f = n.to_string().parse::<f64>().unwrap_or(f64::NAN);
BigFloat::from_f64(f, prec)
}
fn cauchy_bound(poly: &Poly, prec: usize) -> BigFloat {
let rm = RoundingMode::None;
let lc = poly.leading_coeff().cloned().unwrap_or_else(Ratio::one);
if lc.is_zero() {
return BigFloat::from_i32(1, prec);
}
let mut max_ratio = BigFloat::from_i32(0, prec);
for c in poly
.coeffs()
.iter()
.take(poly.coeffs().len().saturating_sub(1))
{
let ratio = c / &lc;
let abs_ratio = if ratio.is_negative() { -ratio } else { ratio };
let bf = ratio_to_bigfloat(&abs_ratio, prec);
if bf.sub(&max_ratio, prec, rm).is_positive() {
max_ratio = bf;
}
}
max_ratio.add(&BigFloat::from_i32(1, prec), prec, rm)
}
fn initial_guesses(poly: &Poly, n: usize, prec: usize, cc: &mut Consts) -> Vec<Complex> {
let rm = RoundingMode::None;
let radius = cauchy_bound(poly, prec);
let center = if n >= 2 && poly.coeffs().len() > n {
let an = &poly.coeffs()[n];
let an1 = &poly.coeffs()[n - 1];
if !an.is_zero() {
let ratio = -(an1 / an) / Ratio::from_integer(BigInt::from(n));
ratio_to_bigfloat(&ratio, prec)
} else {
BigFloat::new(prec)
}
} else {
BigFloat::new(prec)
};
let two_pi = cc.pi(prec, rm).mul(&BigFloat::from_i32(2, prec), prec, rm);
let n_bf = BigFloat::from_i64(n as i64, prec);
let quarter = BigFloat::from_f64(0.25, prec);
let offset = BigFloat::from_f64(0.4, prec);
(0..n)
.map(|k| {
let k_bf = BigFloat::from_i64(k as i64, prec);
let frac = k_bf.add(&quarter, prec, rm).div(&n_bf, prec, rm);
let angle = two_pi.mul(&frac, prec, rm).add(&offset, prec, rm);
let cos_a = angle.cos(prec, rm, cc);
let sin_a = angle.sin(prec, rm, cc);
let re = center.add(&radius.mul(&cos_a, prec, rm), prec, rm);
let im = radius.mul(&sin_a, prec, rm);
(re, im)
})
.collect()
}
pub(crate) fn aberth_roots(poly: &Poly, prec: usize, max_iter: usize) -> Vec<Complex> {
if poly.degree().is_none_or(|d| d == 0) {
return vec![];
}
let rm = RoundingMode::None;
let wp = prec + 64;
let zero_mult = poly.coeffs().iter().take_while(|c| c.is_zero()).count();
let mut roots: Vec<Complex> = (0..zero_mult).map(|_| c_zero(wp)).collect();
let reduced = if zero_mult > 0 {
Poly::from_coeffs(poly.coeffs()[zero_mult..].to_vec())
} else {
poly.clone()
};
let n = match reduced.degree() {
Some(d) if d >= 1 => d,
_ => return roots,
};
let monic = reduced.make_monic();
let deriv = monic.derivative();
let mut cc = match Consts::new() {
Ok(cc) => cc,
Err(e) => {
tracing::warn!(error = ?e, "aberth_roots: astro-float constants init failed");
return roots;
}
};
let mut nonzero = aberth_iterate(&monic, &deriv, n, wp, prec, max_iter, rm, &mut cc);
roots.append(&mut nonzero);
roots.sort_by(|a, b| {
let re_cmp = a.0.cmp(&b.0).unwrap_or(0);
if re_cmp < 0 {
std::cmp::Ordering::Less
} else if re_cmp > 0 {
std::cmp::Ordering::Greater
} else {
let im_cmp = a.1.cmp(&b.1).unwrap_or(0);
if im_cmp < 0 {
std::cmp::Ordering::Less
} else if im_cmp > 0 {
std::cmp::Ordering::Greater
} else {
std::cmp::Ordering::Equal
}
}
});
roots
}
#[allow(clippy::too_many_arguments)]
fn aberth_iterate(
monic: &Poly,
deriv: &Poly,
n: usize,
wp: usize,
prec: usize,
max_iter: usize,
rm: RoundingMode,
cc: &mut Consts,
) -> Vec<Complex> {
let mut roots = initial_guesses(monic, n, wp, cc);
let threshold = BigFloat::from_f64(1e-30, wp);
for _iter in 0..max_iter {
let mut max_correction = BigFloat::new(wp);
let mut corrections: Vec<Complex> = Vec::with_capacity(n);
for i in 0..n {
let p_zi = poly_eval_complex(monic.coeffs(), &roots[i], wp, rm);
let pp_zi = poly_eval_complex(deriv.coeffs(), &roots[i], wp, rm);
let mut sum_recip = c_zero(wp);
for j in 0..n {
if j != i {
let diff = c_sub(&roots[i], &roots[j], wp, rm);
let diff_abs_sq =
diff.0
.mul(&diff.0, wp, rm)
.add(&diff.1.mul(&diff.1, wp, rm), wp, rm);
if diff_abs_sq.is_zero() {
continue;
}
let recip = c_div(&c_one(wp), &diff, wp, rm);
sum_recip = c_add(&sum_recip, &recip, wp, rm);
}
}
let pz_sum = c_mul(&p_zi, &sum_recip, wp, rm);
let denom = c_sub(&pp_zi, &pz_sum, wp, rm);
let denom_abs_sq =
denom
.0
.mul(&denom.0, wp, rm)
.add(&denom.1.mul(&denom.1, wp, rm), wp, rm);
let correction = if denom_abs_sq.is_zero() {
c_zero(wp)
} else {
c_div(&p_zi, &denom, wp, rm)
};
let corr_abs_sq = correction.0.mul(&correction.0, wp, rm).add(
&correction.1.mul(&correction.1, wp, rm),
wp,
rm,
);
if corr_abs_sq.sub(&max_correction, prec, rm).is_positive() {
max_correction = corr_abs_sq;
}
corrections.push(correction);
}
for i in 0..n {
roots[i] = c_sub(&roots[i], &corrections[i], wp, rm);
}
let threshold_sq = threshold.mul(&threshold, wp, rm);
if max_correction.is_zero() || !max_correction.sub(&threshold_sq, prec, rm).is_positive() {
break;
}
}
roots
}
#[allow(dead_code)] pub(crate) fn rootof_eval_f64(poly: &Poly, index: usize) -> Option<(f64, f64)> {
let n = poly.degree()?;
if index >= n {
return None;
}
let roots = aberth_roots(poly, 128, 100);
if index >= roots.len() {
return None;
}
let (re, im) = &roots[index];
let re_f64 = bigfloat_to_f64(re);
let im_f64 = bigfloat_to_f64(im);
Some((re_f64, im_f64))
}
const REAL_AXIS_NOISE: f64 = 1e-6;
fn is_real_root_near(
part: &Poly,
z: (f64, f64),
sturm: &mut Option<super::sturm::SturmChain>,
) -> bool {
let (re, im) = z;
if !re.is_finite() || !im.is_finite() {
return false;
}
let scale = re.abs().max(1.0);
if im.abs() > REAL_AXIS_NOISE * scale {
return false;
}
let Some(center) = crate::base::numeric::f64_to_ratio_exact(re) else {
return false;
};
let Some(eps) = crate::base::numeric::f64_to_ratio_exact(scale * 2f64.powi(-30)) else {
return false;
};
let chain = sturm.get_or_insert_with(|| super::sturm::SturmChain::new(part));
let lo = ¢er - &eps;
let hi = ¢er + &eps;
chain.count_roots_in_closed(&lo, &hi) >= 1
}
fn bigfloat_to_f64(bf: &BigFloat) -> f64 {
let s = format!("{}", bf);
s.parse::<f64>().unwrap_or(f64::NAN)
}
pub(crate) fn nroots_f64(poly: &Poly, prec_bits: usize) -> Vec<(f64, f64)> {
let mut out: Vec<(f64, f64)> = Vec::new();
if poly.degree().unwrap_or(0) == 0 {
return out;
}
let (_content, parts) = poly.sqf_list();
for (part, mult) in parts {
if part.degree().unwrap_or(0) == 0 {
continue;
}
let max_iter = 100 + 20 * part.degree().unwrap_or(0);
let roots = aberth_roots(&part, prec_bits, max_iter);
let mut sturm: Option<super::sturm::SturmChain> = None;
for (re, im) in roots {
let mut pair = (bigfloat_to_f64(&re), bigfloat_to_f64(&im));
if pair.1 != 0.0 && is_real_root_near(&part, pair, &mut sturm) {
pair.1 = 0.0;
}
for _ in 0..mult {
out.push(pair);
}
}
}
out.sort_by(|a, b| {
a.0.partial_cmp(&b.0)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal))
});
out
}
#[cfg(test)]
mod tests {
use super::*;
fn poly_from_coeffs(coeffs: &[i64]) -> Poly {
let rat_coeffs: Vec<Ratio<BigInt>> = coeffs
.iter()
.map(|&c| Ratio::from_integer(BigInt::from(c)))
.collect();
Poly::from_coeffs(rat_coeffs)
}
#[test]
fn aberth_quadratic_real_roots() {
let poly = poly_from_coeffs(&[6, -5, 1]);
let roots = aberth_roots(&poly, 128, 100);
assert_eq!(roots.len(), 2);
let mut real_parts: Vec<f64> = roots.iter().map(|r| bigfloat_to_f64(&r.0)).collect();
real_parts.sort_by(|a, b| a.partial_cmp(b).unwrap());
assert!(
(real_parts[0] - 2.0).abs() < 1e-10,
"root 0: {}",
real_parts[0]
);
assert!(
(real_parts[1] - 3.0).abs() < 1e-10,
"root 1: {}",
real_parts[1]
);
}
#[test]
fn aberth_quadratic_complex_roots() {
let poly = poly_from_coeffs(&[1, 0, 1]);
let roots = aberth_roots(&poly, 128, 100);
assert_eq!(roots.len(), 2);
for root in &roots {
let re = bigfloat_to_f64(&root.0);
let im = bigfloat_to_f64(&root.1);
assert!(re.abs() < 1e-10, "real part should be ~0: {re}");
assert!((im.abs() - 1.0).abs() < 1e-10, "|im| should be ~1: {im}");
}
}
#[test]
fn aberth_quintic() {
let poly = poly_from_coeffs(&[-1, -1, 0, 0, 0, 1]);
let roots = aberth_roots(&poly, 128, 200);
assert_eq!(roots.len(), 5);
for (i, root) in roots.iter().enumerate() {
let val = poly_eval_complex(poly.coeffs(), root, 128, RoundingMode::None);
let mag_sq = bigfloat_to_f64(&val.0).powi(2) + bigfloat_to_f64(&val.1).powi(2);
assert!(
mag_sq < 1e-15,
"root {i} residual too large: |p(z)|^2 = {mag_sq}"
);
}
let real_roots: Vec<_> = roots
.iter()
.filter(|r| bigfloat_to_f64(&r.1).abs() < 1e-8)
.collect();
assert_eq!(real_roots.len(), 1, "should have exactly 1 real root");
let real_val = bigfloat_to_f64(&real_roots[0].0);
assert!(
(real_val - 1.1673).abs() < 0.001,
"real root ≈ 1.1673, got {real_val}"
);
}
#[test]
fn rootof_eval_f64_basic() {
let poly = poly_from_coeffs(&[-4, 0, 1]);
let r0 = rootof_eval_f64(&poly, 0).unwrap();
let r1 = rootof_eval_f64(&poly, 1).unwrap();
assert!(r0.1.abs() < 1e-10);
assert!(r1.1.abs() < 1e-10);
let mut reals = [r0.0, r1.0];
reals.sort_by(|a, b| a.partial_cmp(b).unwrap());
assert!((reals[0] - (-2.0)).abs() < 1e-10);
assert!((reals[1] - 2.0).abs() < 1e-10);
}
#[test]
fn rootof_out_of_range() {
let poly = poly_from_coeffs(&[-1, 0, 1]); assert!(rootof_eval_f64(&poly, 2).is_none());
assert!(rootof_eval_f64(&poly, 100).is_none());
}
#[test]
fn nroots_with_multiplicity() {
let poly = poly_from_coeffs(&[2, -3, 0, 1]);
let roots = nroots_f64(&poly, 128);
assert_eq!(roots.len(), 3);
assert!((roots[0].0 + 2.0).abs() < 1e-12, "{roots:?}");
assert!((roots[1].0 - 1.0).abs() < 1e-12, "{roots:?}");
assert!((roots[2].0 - 1.0).abs() < 1e-12, "{roots:?}");
assert!(roots.iter().all(|r| r.1.abs() < 1e-12));
}
#[test]
fn nroots_wilkinson_like_degree_10() {
let mut poly = poly_from_coeffs(&[1]);
for k in 1..=10 {
poly = &poly * &poly_from_coeffs(&[-k, 1]);
}
let roots = nroots_f64(&poly, 192);
assert_eq!(roots.len(), 10);
for (i, r) in roots.iter().enumerate() {
assert!((r.0 - (i as f64 + 1.0)).abs() < 1e-8, "root {i}: {r:?}");
assert!(r.1.abs() < 1e-8);
}
}
}