use super::softfloat::{self, ExcFlags, RoundCtx};
use super::types::FloatX80;
use std::sync::OnceLock;
#[inline]
fn fadd(a: FloatX80, b: FloatX80) -> FloatX80 {
softfloat::add(a, b, RoundCtx::NEAREST_EXT, &mut ExcFlags::default())
}
#[inline]
fn fsub(a: FloatX80, b: FloatX80) -> FloatX80 {
softfloat::sub(a, b, RoundCtx::NEAREST_EXT, &mut ExcFlags::default())
}
#[inline]
fn fmul(a: FloatX80, b: FloatX80) -> FloatX80 {
softfloat::mul(a, b, RoundCtx::NEAREST_EXT, &mut ExcFlags::default())
}
#[inline]
fn fdiv(a: FloatX80, b: FloatX80) -> FloatX80 {
softfloat::div(a, b, RoundCtx::NEAREST_EXT, &mut ExcFlags::default())
}
#[inline]
fn fx(v: f64) -> FloatX80 {
FloatX80::from_f64(v)
}
#[inline]
fn two_sum(a: FloatX80, b: FloatX80) -> (FloatX80, FloatX80) {
let s = fadd(a, b);
let bb = fsub(s, a);
let err = fadd(fsub(a, fsub(s, bb)), fsub(b, bb));
(s, err)
}
#[inline]
fn quick_two_sum(a: FloatX80, b: FloatX80) -> (FloatX80, FloatX80) {
let s = fadd(a, b);
let e = fsub(b, fsub(s, a));
(s, e)
}
#[inline]
fn split(a: FloatX80) -> (FloatX80, FloatX80) {
const SPLIT: f64 = 4294967297.0; let t = fmul(fx(SPLIT), a);
let hi = fsub(t, fsub(t, a));
let lo = fsub(a, hi);
(hi, lo)
}
#[inline]
fn two_prod(a: FloatX80, b: FloatX80) -> (FloatX80, FloatX80) {
let p = fmul(a, b);
let (ah, al) = split(a);
let (bh, bl) = split(b);
let e = fadd(
fadd(fadd(fsub(fmul(ah, bh), p), fmul(ah, bl)), fmul(al, bh)),
fmul(al, bl),
);
(p, e)
}
#[derive(Clone, Copy)]
pub struct Df {
pub hi: FloatX80,
pub lo: FloatX80,
}
impl Df {
#[inline]
pub fn from_x80(a: FloatX80) -> Df {
Df {
hi: a,
lo: FloatX80::zero(false),
}
}
#[inline]
pub fn from_i32(n: i32) -> Df {
Df::from_x80(softfloat::from_u64(n.unsigned_abs() as u64, n < 0))
}
#[inline]
pub fn neg(self) -> Df {
Df {
hi: softfloat::neg(self.hi),
lo: softfloat::neg(self.lo),
}
}
#[inline]
pub fn to_x80(self, ctx: RoundCtx, f: &mut ExcFlags) -> FloatX80 {
softfloat::add(self.hi, self.lo, ctx, f)
}
}
#[inline]
pub fn add(a: Df, b: Df) -> Df {
let (s, e) = two_sum(a.hi, b.hi);
let e = fadd(e, fadd(a.lo, b.lo));
let (hi, lo) = quick_two_sum(s, e);
Df { hi, lo }
}
#[inline]
pub fn sub(a: Df, b: Df) -> Df {
add(a, b.neg())
}
#[inline]
pub fn add_x80(a: Df, b: FloatX80) -> Df {
let (s, e) = two_sum(a.hi, b);
let e = fadd(e, a.lo);
let (hi, lo) = quick_two_sum(s, e);
Df { hi, lo }
}
#[inline]
pub fn mul(a: Df, b: Df) -> Df {
let (p, e) = two_prod(a.hi, b.hi);
let e = fadd(e, fadd(fmul(a.hi, b.lo), fmul(a.lo, b.hi)));
let (hi, lo) = quick_two_sum(p, e);
Df { hi, lo }
}
#[inline]
pub fn mul_x80(a: Df, b: FloatX80) -> Df {
let (p, e) = two_prod(a.hi, b);
let e = fadd(e, fmul(a.lo, b));
let (hi, lo) = quick_two_sum(p, e);
Df { hi, lo }
}
#[inline]
pub fn sqr(a: Df) -> Df {
mul(a, a)
}
pub fn div(a: Df, b: Df) -> Df {
let q1 = fdiv(a.hi, b.hi);
let r = sub(a, mul_x80(b, q1));
let q2 = fdiv(r.hi, b.hi);
let r = sub(r, mul_x80(b, q2));
let q3 = fdiv(r.hi, b.hi);
let (hi, lo) = quick_two_sum(q1, q2);
add_x80(Df { hi, lo }, q3)
}
#[inline]
pub fn recip(b: Df) -> Df {
div(Df::from_i32(1), b)
}
fn negligible(pow: FloatX80, sum: FloatX80) -> bool {
if pow.is_zero() {
return true;
}
if sum.is_zero() {
return false;
}
(sum.biased_exp() as i32) - (pow.biased_exp() as i32) > 130
}
fn odd_series(x: Df, alternating: bool) -> Df {
let x2 = sqr(x);
let mut pow = x; let mut sum = x;
let mut k = 1usize;
loop {
pow = mul(pow, x2);
let term = div(pow, Df::from_i32((2 * k + 1) as i32));
let term = if alternating && (k & 1 == 1) {
term.neg()
} else {
term
};
sum = add(sum, term);
if negligible(pow.hi, sum.hi) || k > 200 {
break;
}
k += 1;
}
sum
}
pub fn term_negligible(term: Df, sum: Df) -> bool {
negligible(term.hi, sum.hi)
}
pub fn atan_small(x: Df) -> Df {
odd_series(x, true)
}
pub fn atanh_small(x: Df) -> Df {
odd_series(x, false)
}
#[derive(Clone, Copy)]
pub struct Consts {
pub pi: Df,
pub pi_2: Df,
pub ln2: Df,
pub ln10: Df,
pub log2e: Df,
pub log10e: Df,
}
static CONSTS: OnceLock<Consts> = OnceLock::new();
pub fn consts() -> &'static Consts {
CONSTS.get_or_init(|| {
let inv = |n: i32| div(Df::from_i32(1), Df::from_i32(n));
let pi = sub(
mul_x80(atan_small(inv(5)), fx(16.0)),
mul_x80(atan_small(inv(239)), fx(4.0)),
);
let ln2 = mul_x80(atanh_small(inv(3)), fx(2.0));
let ln10 = add(mul_x80(ln2, fx(3.0)), mul_x80(atanh_small(inv(9)), fx(2.0)));
Consts {
pi,
pi_2: mul_x80(pi, fx(0.5)),
ln2,
ln10,
log2e: recip(ln2),
log10e: recip(ln10),
}
})
}
#[cfg(test)]
mod tests {
use super::*;
fn rn() -> RoundCtx {
RoundCtx::NEAREST_EXT
}
fn to_f64(d: Df) -> f64 {
d.to_x80(rn(), &mut ExcFlags::default()).to_f64()
}
#[test]
fn two_prod_is_exact() {
let a = fx(1.5);
let b = fx(3.25);
let (p, e) = two_prod(a, b);
assert_eq!(fadd(p, e).to_f64(), 1.5 * 3.25);
assert!(e.is_zero()); }
#[test]
fn df_arithmetic_matches_f64() {
let a = Df::from_i32(7);
let b = Df::from_i32(3);
assert_eq!(to_f64(add(a, b)), 10.0);
assert_eq!(to_f64(sub(a, b)), 4.0);
assert_eq!(to_f64(mul(a, b)), 21.0);
assert!((to_f64(div(a, b)) - 7.0 / 3.0).abs() < 1e-15);
}
#[test]
fn df_keeps_more_than_64_bits() {
let tiny = FloatX80 {
sign_exp: (16383 - 80) as u16,
mantissa: 0x8000_0000_0000_0000,
};
let d = add_x80(Df::from_i32(1), tiny);
assert_eq!(d.hi.to_f64(), 1.0);
assert!(!d.lo.is_zero());
}
#[test]
fn constants_match_rom_hi_words() {
let c = consts();
assert_eq!(c.pi.hi, softfloat::const_rom(0x00));
assert_eq!(c.ln2.hi, softfloat::const_rom(0x30));
assert_eq!(c.ln10.hi, softfloat::const_rom(0x31));
assert_eq!(c.log2e.hi, softfloat::const_rom(0x0D));
assert_eq!(c.log10e.hi, softfloat::const_rom(0x0E));
}
#[test]
fn constants_value_sanity() {
let c = consts();
assert!((to_f64(c.pi) - std::f64::consts::PI).abs() < 1e-15);
assert!((to_f64(c.ln2) - std::f64::consts::LN_2).abs() < 1e-15);
assert!((to_f64(c.ln10) - std::f64::consts::LN_10).abs() < 1e-15);
}
}