use thermite::{
element::FloatElementWithBits,
mask::GenericMask,
math::{CoreMathWithPolicy as _, TranscendentalMathWithPolicy as _, policy::Policy, specialized::FlushDenormals},
prelude::*,
};
use crate::specialized::SpecializedSpecialMath;
use crate::tables::Digamma;
#[inline(always)]
pub fn digamma_impl<P, E, V, const NR: usize, const NL: usize, const NP: usize, const NQ: usize>(
x_in: V,
t: &Digamma<E, NR, NL, NP, NQ>,
) -> V
where
P: Policy,
E: FloatElementWithBits,
V: FloatVectorWithBits<Element = E> + SpecializedSpecialMath<E>,
{
let mut x0 = x_in;
#[cfg(not(target_arch = "spirv"))]
if let Some(new_x) = FlushDenormals::<P>::flush_denormals([x0]) {
x0 = new_x[0];
}
let mut result = V::ZERO;
let mut x = x0;
let reflect = x0.cmp_le(V::NEG_ONE);
let mut refl_pole = GenericMask::FALSY;
if const { P::POLICY.avoid_branching } || reflect.any() {
let xr = V::ONE - x0; let mut rem = xr - xr.floor();
rem = rem.sub_c(rem.cmp_gt(V::HALF), V::ONE);
let refl_term = V::PI / rem.tan_pi_p::<P>();
result = refl_term.zz(reflect); x = reflect.select(xr, x);
refl_pole = reflect & rem.is_zero(); }
let large = x.cmp_ge(V::splat(E::from_int(10)));
let mut active = (x.cmp_gt(V::TWO) | x.cmp_lt(V::ONE)) & !large;
while active.any() {
let sign = (x - V::ONE).signum();
let xs = x.sub_c(active, sign); let term = sign * x.min(xs).reciprocal_p::<P>();
result = result.add_c(active, term);
x = xs;
active = (x.cmp_gt(V::TWO) | x.cmp_lt(V::ONE)) & !large;
}
let xm1 = x - V::ONE;
let mut g = x;
let mut i = 0;
while i < NR {
g -= V::splat(t.roots[i]);
i += 1;
}
let r = xm1.poly_p::<P, _>(&t.p_12) / xm1.poly_p::<P, _>(&t.q_12);
let rational = g * (V::splat(t.y) + r);
let z = (xm1 * xm1).reciprocal_p::<P>();
let asymptotic = z.nmul_adde(
z.poly_p::<P, _>(&t.p_large),
xm1.ln_p::<P>() + (xm1 + xm1).reciprocal_p::<P>(),
);
let mut res = result + large.select(asymptotic, rational);
if const { P::POLICY.check_overflow } {
let pole = x0.is_zero() | refl_pole;
res = pole.select(V::NAN, res);
res = x0.is_nan().select(V::NAN, res);
}
res
}