use crate::common::consts::ONE;
use crate::common::util::calc_add_cost;
use crate::common::util::calc_mul_cost;
use crate::common::util::log2_floor;
use crate::common::util::sqrt_int;
use crate::defs::Error;
use crate::defs::RoundingMode;
use crate::num::ExactNumNumber;
use crate::Sign;
#[cfg(not(feature = "std"))]
use alloc::vec::Vec;
const MAX_CACHE: usize = 128;
const RECT_ITER_THRESHOLD: usize = MAX_CACHE / 10 * 9;
pub(crate) trait PolycoeffGen {
fn next(&mut self, rm: RoundingMode) -> Result<&ExactNumNumber, Error>;
fn iter_cost(&self) -> usize;
fn is_div(&self) -> bool {
false
}
}
pub(crate) struct FactPolycoeffGen {
one_full_p: ExactNumNumber,
inc: ExactNumNumber,
fct: ExactNumNumber,
sign: i8,
flip_sign: bool,
iter_cost: usize,
}
impl FactPolycoeffGen {
pub(crate) fn for_sin(p: usize) -> Result<Self, Error> {
Self::new(p, true, true)
}
pub(crate) fn for_cos(p: usize) -> Result<Self, Error> {
Self::new(p, false, true)
}
pub(crate) fn for_sinh(p: usize) -> Result<Self, Error> {
Self::new(p, true, false)
}
fn new(p: usize, start_from_one: bool, flip_sign: bool) -> Result<Self, Error> {
let inc = if start_from_one {
ExactNumNumber::from_word(1, 1)?
} else {
ExactNumNumber::new(1)?
};
let fct = ExactNumNumber::from_word(1, p)?;
let one_full_p = ExactNumNumber::from_word(1, p)?;
let iter_cost =
(calc_mul_cost(p) + calc_add_cost(p) + calc_add_cost(inc.mantissa_max_bit_len())) * 2;
Ok(Self {
one_full_p,
inc,
fct,
sign: 1,
flip_sign,
iter_cost,
})
}
}
impl PolycoeffGen for FactPolycoeffGen {
fn next(&mut self, rm: RoundingMode) -> Result<&ExactNumNumber, Error> {
let p_inc = self.inc.mantissa_max_bit_len();
let p_one = self.one_full_p.mantissa_max_bit_len();
self.inc = self.inc.add(&ONE, p_inc, rm)?;
let inv_inc = self.one_full_p.div(&self.inc, p_one, rm)?;
self.fct = self.fct.mul(&inv_inc, p_one, rm)?;
self.inc = self.inc.add(&ONE, p_inc, rm)?;
let inv_inc = self.one_full_p.div(&self.inc, p_one, rm)?;
self.fct = self.fct.mul(&inv_inc, p_one, rm)?;
if self.flip_sign {
self.sign *= -1;
if self.sign > 0 {
self.fct.set_sign(Sign::Pos);
} else {
self.fct.set_sign(Sign::Neg);
}
}
Ok(&self.fct)
}
#[inline]
fn iter_cost(&self) -> usize {
self.iter_cost
}
}
pub trait ArgReductionEstimator {
fn reduction_cost(n: usize, p: usize) -> u64;
fn reduction_effect(n: usize, m: isize) -> usize;
}
pub(crate) fn series_cost_optimize<S: ArgReductionEstimator>(
p: usize,
polycoeff_gen: &impl PolycoeffGen,
m: isize,
pwr_step: usize,
ext: bool,
) -> (usize, usize, usize) {
let reduction_num_step = log2_floor(p) / 2;
if reduction_num_step == 0 {
let m_eff = S::reduction_effect(0, m);
let niter = series_niter(p, m_eff) / pwr_step.max(1);
return (0, niter, m_eff);
}
let mut reduction_times = if reduction_num_step as isize > m {
(reduction_num_step as isize - m) as usize
} else {
0
};
let mut cost1 = u64::MAX;
loop {
let m_eff = S::reduction_effect(reduction_times, m);
let niter = series_niter(p, m_eff) / pwr_step.max(1);
let cost2 = if ext {
polycoeff_gen.iter_cost() as u64 * niter as u64
} else {
series_cost(niter, p, polycoeff_gen)
} + S::reduction_cost(reduction_times, p);
if cost2 < cost1 {
cost1 = cost2;
reduction_times += reduction_num_step;
if reduction_times > p.saturating_add(reduction_num_step) {
return (reduction_times - reduction_num_step, niter, m_eff);
}
} else {
return (reduction_times - reduction_num_step, niter, m_eff);
}
}
}
pub(crate) fn series_run<T: PolycoeffGen>(
acc: ExactNumNumber,
x_first: ExactNumNumber,
x_step: ExactNumNumber,
niter: usize,
polycoeff_gen: &mut T,
) -> Result<ExactNumNumber, Error> {
let mut ret = if x_first.is_zero() || x_step.is_zero() {
series_compute_fast(acc, x_first, polycoeff_gen)
} else if niter >= RECT_ITER_THRESHOLD {
series_rectangular(niter, acc, x_first, x_step, polycoeff_gen)
} else if polycoeff_gen.is_div() {
series_linear(acc, x_first, x_step, polycoeff_gen)
} else {
series_horner(acc, x_first, x_step, polycoeff_gen)
}?;
ret.set_inexact(true);
Ok(ret)
}
fn series_niter(p: usize, m: usize) -> usize {
let ln = log2_floor(p);
let lln = log2_floor(ln);
p / (ln - lln + m - 2)
}
fn series_cost<T: PolycoeffGen>(niter: usize, p: usize, polycoeff_gen: &T) -> u64 {
let cost_mul = calc_mul_cost(p);
let cost_add = calc_add_cost(p);
let cost = niter as u64 * (cost_mul + cost_add + polycoeff_gen.iter_cost()) as u64;
if niter >= RECT_ITER_THRESHOLD {
cost + sqrt_int(niter as u32) as u64 * cost_mul as u64
+ niter as u64 / 10 * ((cost_mul << 1) + cost_add + polycoeff_gen.iter_cost()) as u64
} else {
cost
}
}
fn series_compute_fast<T: PolycoeffGen>(
acc: ExactNumNumber,
x_first: ExactNumNumber,
polycoeff_gen: &mut T,
) -> Result<ExactNumNumber, Error> {
if x_first.is_zero() {
Ok(acc)
} else {
let p = acc
.mantissa_max_bit_len()
.max(x_first.mantissa_max_bit_len());
let is_div = polycoeff_gen.is_div();
let coeff = polycoeff_gen.next(RoundingMode::None)?;
let part = if is_div {
x_first.div(coeff, p, RoundingMode::None)
} else {
x_first.mul(coeff, p, RoundingMode::None)
}?;
acc.add(&part, p, RoundingMode::None)
}
}
fn series_rectangular<T: PolycoeffGen>(
mut niter: usize,
add: ExactNumNumber,
x_first: ExactNumNumber,
x_step: ExactNumNumber,
polycoeff_gen: &mut T,
) -> Result<ExactNumNumber, Error> {
debug_assert!(niter >= 4);
let p = add
.mantissa_max_bit_len()
.max(x_first.mantissa_max_bit_len())
.max(x_step.mantissa_max_bit_len());
let mut acc = ExactNumNumber::new(p)?;
let mut cache = Vec::<ExactNumNumber>::new();
let sqrt_iter = sqrt_int(niter as u32) as usize;
let cache_sz = MAX_CACHE.min(sqrt_iter);
cache.try_reserve_exact(cache_sz)?;
let mut x_pow = x_step.clone()?;
for _ in 0..cache_sz {
cache.push(x_pow.clone()?);
x_pow = x_pow.mul(&x_step, p, RoundingMode::None)?;
}
let poly_val = compute_row(p, &cache, polycoeff_gen)?;
acc = acc.add(&poly_val, p, RoundingMode::None)?;
let mut terminal_pow = x_pow.clone()?;
niter -= cache_sz;
loop {
let poly_val = compute_row(p, &cache, polycoeff_gen)?;
let part = poly_val.mul(&terminal_pow, p, RoundingMode::None)?;
acc = acc.add(&part, p, RoundingMode::None)?;
terminal_pow = terminal_pow.mul(&x_pow, p, RoundingMode::None)?;
niter -= cache_sz;
if niter < cache_sz {
break;
}
}
drop(cache);
acc = acc.mul(&x_first, p, RoundingMode::None)?;
terminal_pow = terminal_pow.mul(&x_first, p, RoundingMode::None)?;
acc = acc.add(&add, p, RoundingMode::None)?;
acc = if niter < MAX_CACHE * 10 && !polycoeff_gen.is_div() {
series_horner(acc, terminal_pow, x_step, polycoeff_gen)
} else {
series_linear(acc, terminal_pow, x_step, polycoeff_gen)
}?;
Ok(acc)
}
fn series_linear<T: PolycoeffGen>(
mut acc: ExactNumNumber,
x_first: ExactNumNumber,
x_step: ExactNumNumber,
polycoeff_gen: &mut T,
) -> Result<ExactNumNumber, Error> {
let p = acc
.mantissa_max_bit_len()
.max(x_first.mantissa_max_bit_len())
.max(x_step.mantissa_max_bit_len());
let is_div = polycoeff_gen.is_div();
let mut x_pow = x_first;
loop {
let coeff = polycoeff_gen.next(RoundingMode::None)?;
let part = if is_div {
x_pow.div(coeff, p, RoundingMode::None)
} else {
x_pow.mul(coeff, p, RoundingMode::None)
}?;
acc = acc.add(&part, p, RoundingMode::None)?;
if part.exponent() as isize <= acc.exponent() as isize - acc.mantissa_max_bit_len() as isize
{
break;
}
x_pow = x_pow.mul(&x_step, p, RoundingMode::None)?;
}
Ok(acc)
}
fn compute_row<T: PolycoeffGen>(
p: usize,
cache: &[ExactNumNumber],
polycoeff_gen: &mut T,
) -> Result<ExactNumNumber, Error> {
let is_div = polycoeff_gen.is_div();
let mut acc = ExactNumNumber::new(p)?;
let coeff = polycoeff_gen.next(RoundingMode::None)?;
if is_div {
let r = coeff.reciprocal(p, RoundingMode::None)?;
acc = acc.add(&r, p, RoundingMode::None)?;
} else {
acc = acc.add(coeff, p, RoundingMode::None)?;
}
for x_pow in cache {
let coeff = polycoeff_gen.next(RoundingMode::None)?;
let add = if is_div {
x_pow.div(coeff, p, RoundingMode::None)
} else {
x_pow.mul(coeff, p, RoundingMode::None)
}?;
acc = acc.add(&add, p, RoundingMode::None)?;
}
Ok(acc)
}
fn series_horner<T: PolycoeffGen>(
add: ExactNumNumber,
x_first: ExactNumNumber,
x_step: ExactNumNumber,
polycoeff_gen: &mut T,
) -> Result<ExactNumNumber, Error> {
debug_assert!(x_first.exponent() <= 0);
debug_assert!(x_step.exponent() <= 0);
debug_assert!(!polycoeff_gen.is_div());
let p = add
.mantissa_max_bit_len()
.max(x_first.mantissa_max_bit_len())
.max(x_step.mantissa_max_bit_len());
let mut cache = Vec::<ExactNumNumber>::new();
let mut x_p = -(x_first.exponent() as isize) - x_step.exponent() as isize;
let mut coef_p = 0;
while x_p + coef_p < p as isize - add.exponent() as isize {
let coeff = polycoeff_gen.next(RoundingMode::None)?;
coef_p = -coeff.exponent() as isize;
x_p += -x_step.exponent() as isize;
cache.push(coeff.clone()?);
}
let last_coeff = polycoeff_gen.next(RoundingMode::None)?;
let mut acc = last_coeff.clone()?;
for coeff in cache.iter().rev() {
acc = acc.mul(&x_step, p, RoundingMode::None)?;
acc = acc.add(coeff, p, RoundingMode::None)?;
}
acc = acc.mul(&x_first, p, RoundingMode::None)?;
add.add(&acc, p, RoundingMode::None)
}