use num_traits::Zero;
use rug::Integer;
use std::ops::{BitAnd, BitOr};
use crate::ieee754::{Exceptions, IEEE754Val, IEEE754};
use crate::rfloat::{RFloat, RFloatContext};
use crate::util::bitmask;
use crate::{Real, RoundingContext, RoundingDirection, RoundingMode, Split};
#[derive(Clone, Debug)]
pub struct IEEE754Context {
es: usize,
nbits: usize,
rm: RoundingMode,
ftz: bool,
}
impl IEEE754Context {
pub const ES_MAX: usize = 32;
pub const ES_MIN: usize = 2;
pub const PREC_MIN: usize = 3;
pub fn new(es: usize, nbits: usize) -> Self {
assert!(
es >= Self::ES_MIN,
"exponent width needs to be at least {} bits, given {} bits",
Self::ES_MIN,
es
);
assert!(
es <= Self::ES_MAX,
"exponent width needs to be at most {} bits, given {} bits",
Self::ES_MAX,
es
);
assert!(
nbits >= es + Self::PREC_MIN,
"total bitwidth needs to be at least {} bits, given {} bits",
es + Self::PREC_MIN,
nbits
);
Self {
es,
nbits,
rm: RoundingMode::NearestTiesToEven,
ftz: false,
}
}
pub fn with_rounding_mode(mut self, rm: RoundingMode) -> Self {
self.rm = rm;
self
}
pub fn with_ftz(mut self, enable: bool) -> Self {
self.ftz = enable;
self
}
pub fn es(&self) -> usize {
self.es
}
pub fn rm(&self) -> RoundingMode {
self.rm
}
pub fn ftz(&self) -> bool {
self.ftz
}
pub fn nbits(&self) -> usize {
self.nbits
}
pub fn max_p(&self) -> usize {
self.nbits - self.es
}
pub fn max_m(&self) -> usize {
self.nbits - self.es - 1
}
pub fn emax(&self) -> isize {
(1 << (self.es - 1)) - 1
}
pub fn emin(&self) -> isize {
1 - self.emax()
}
pub fn expmax(&self) -> isize {
self.emax() - (self.max_m() as isize)
}
pub fn expmin(&self) -> isize {
self.emin() - (self.max_m() as isize)
}
pub fn bias(&self) -> isize {
self.emax()
}
pub fn zero(&self, sign: bool) -> IEEE754 {
IEEE754 {
num: if sign {
IEEE754Val::NegZero
} else {
IEEE754Val::PosZero
},
flags: Exceptions::default(),
ctx: self.clone(),
}
}
pub fn min_float(&self, sign: bool) -> IEEE754 {
IEEE754 {
num: IEEE754Val::Subnormal(sign, Integer::from(1)),
flags: Exceptions::default(),
ctx: self.clone(),
}
}
pub fn max_float(&self, sign: bool) -> IEEE754 {
IEEE754 {
num: IEEE754Val::Normal(sign, self.expmax(), bitmask(self.max_p())),
flags: Exceptions::default(),
ctx: self.clone(),
}
}
pub fn inf(&self, sign: bool) -> IEEE754 {
IEEE754 {
num: if sign {
IEEE754Val::NegInfinity
} else {
IEEE754Val::PosInfinity
},
flags: Default::default(),
ctx: self.clone(),
}
}
pub fn qnan(&self) -> IEEE754 {
IEEE754 {
num: IEEE754Val::Nan(false, true, Integer::from(0)),
flags: Default::default(),
ctx: self.clone(),
}
}
pub fn snan(&self) -> IEEE754 {
IEEE754 {
num: IEEE754Val::Nan(false, false, Integer::from(1)),
flags: Default::default(),
ctx: self.clone(),
}
}
pub fn bits_to_number(&self, b: Integer) -> IEEE754 {
let p = self.nbits - self.es;
let limit = Integer::from(1) << self.nbits;
assert!(b < limit, "must be less than 1 << nbits");
let s = b.get_bit((self.nbits - 1) as u32);
let e = (b.clone() >> (p - 1)).bitand(bitmask(self.es));
let m = b.bitand(bitmask(p - 1));
let e_norm = e.to_isize().unwrap() - self.emax();
let num = if e_norm < self.emin() {
if m.is_zero() {
if s {
IEEE754Val::NegZero
} else {
IEEE754Val::PosZero
}
} else {
IEEE754Val::Subnormal(s, m)
}
} else if e_norm <= self.emax() {
let c = (Integer::from(1) << (p - 1)).bitor(m);
let exp = e_norm - (p as isize - 1);
IEEE754Val::Normal(s, exp, c)
} else {
if m.is_zero() {
if s {
IEEE754Val::NegInfinity
} else {
IEEE754Val::PosInfinity
}
} else {
let quiet = m.get_bit((p - 2) as u32);
let payload = m.bitand(bitmask(p - 2));
IEEE754Val::Nan(s, quiet, payload)
}
};
IEEE754 {
num,
flags: Exceptions::default(),
ctx: self.clone(),
}
}
}
impl IEEE754Context {
fn overflow_to_infinity(sign: bool, rm: RoundingMode) -> bool {
match rm.to_direction(sign) {
(true, _) => true,
(_, RoundingDirection::ToZero) => false, (_, RoundingDirection::AwayZero) => true, (_, RoundingDirection::ToEven) => true, (_, RoundingDirection::ToOdd) => false, }
}
fn round_tiny<T: Real>(&self, num: &T) -> bool {
if num.is_zero() {
return false;
}
let e_trunc = num.e().unwrap();
match e_trunc.cmp(&(self.emin() - 1)) {
std::cmp::Ordering::Less => {
true
}
std::cmp::Ordering::Greater => {
false
}
std::cmp::Ordering::Equal => {
let unbounded_ctx = RFloatContext::new()
.with_rounding_mode(self.rm)
.with_max_p(self.max_p());
let unbounded = unbounded_ctx.round(num);
unbounded.e().unwrap() < self.emin()
}
}
}
fn round_finalize(
&self,
unbounded: RFloat,
tiny_pre: bool,
tiny_post: bool,
inexact: bool,
carry: bool,
) -> IEEE754 {
let sign = unbounded.sign().unwrap();
if unbounded.is_zero() {
return IEEE754 {
num: if sign {
IEEE754Val::NegZero
} else {
IEEE754Val::PosZero
},
flags: Exceptions {
underflow_pre: tiny_pre && inexact,
underflow_post: tiny_post && inexact,
inexact,
tiny_pre,
tiny_post,
..Default::default()
},
ctx: self.clone(),
};
}
let e = unbounded.e().unwrap();
if e > self.emax() {
if IEEE754Context::overflow_to_infinity(sign, self.rm) {
return IEEE754 {
num: if sign {
IEEE754Val::NegInfinity
} else {
IEEE754Val::PosInfinity
},
flags: Exceptions {
overflow: true,
inexact: true,
..Default::default()
},
ctx: self.clone(),
};
} else {
let mut maxfloat = self.max_float(sign);
maxfloat.flags.overflow = true;
maxfloat.flags.inexact = true;
return maxfloat;
}
}
if self.ftz && tiny_post {
return IEEE754 {
num: if sign {
IEEE754Val::NegZero
} else {
IEEE754Val::PosZero
},
flags: Exceptions {
underflow_pre: true,
underflow_post: true,
inexact: true,
tiny_pre: true,
tiny_post: true,
..Default::default()
},
ctx: self.clone(),
};
}
let c = unbounded.c().unwrap();
if e < self.emin() {
IEEE754 {
num: IEEE754Val::Subnormal(sign, c),
flags: Exceptions {
underflow_pre: tiny_pre && inexact,
underflow_post: tiny_post && inexact,
inexact,
tiny_pre,
tiny_post,
..Default::default()
},
ctx: self.clone(),
}
} else {
let exp = unbounded.exp().unwrap();
IEEE754 {
num: IEEE754Val::Normal(sign, exp, c),
flags: Exceptions {
underflow_pre: tiny_pre && inexact,
underflow_post: tiny_post && inexact,
inexact,
carry,
tiny_pre,
tiny_post,
..Default::default()
},
ctx: self.clone(),
}
}
}
}
impl RoundingContext for IEEE754Context {
type Format = IEEE754;
fn round<T: Real>(&self, num: &T) -> Self::Format {
if num.is_zero() {
IEEE754 {
num: match num.sign() {
Some(true) => IEEE754Val::NegZero,
_ => IEEE754Val::PosZero,
},
flags: Exceptions::default(),
ctx: self.clone(),
}
} else if num.is_infinite() {
IEEE754 {
num: match num.sign() {
Some(true) => IEEE754Val::NegInfinity,
_ => IEEE754Val::PosInfinity,
},
flags: Exceptions::default(),
ctx: self.clone(),
}
} else if num.is_nar() {
let sign = num.sign().unwrap_or(false);
IEEE754 {
num: IEEE754Val::Nan(sign, true, Integer::zero()),
flags: Exceptions::default(),
ctx: self.clone(),
}
} else {
let (p, n) = RFloatContext::new()
.with_max_p(self.max_p())
.with_min_n(self.expmin() - 1)
.round_params(num);
let split = Split::new(num, p, n);
let inexact = !split.is_exact();
let unrounded_e = split.e();
let (tiny_pre, tiny_post) = match unrounded_e {
None => (false, false), Some(e) => {
let tiny_pre = e < self.emin();
let tiny_post = self.round_tiny(&split);
(tiny_pre, tiny_post)
}
};
let unbounded = RFloatContext::round_finalize(split, self.rm);
let carry = match (unrounded_e, unbounded.e()) {
(Some(e1), Some(e2)) => e2 > e1,
(_, _) => false,
};
self.round_finalize(unbounded, tiny_pre, tiny_post, inexact, carry)
}
}
}