diff options
Diffstat (limited to 'src/eval.rs')
| -rw-r--r-- | src/eval.rs | 380 |
1 files changed, 247 insertions, 133 deletions
diff --git a/src/eval.rs b/src/eval.rs index e40743b..d9fb184 100644 --- a/src/eval.rs +++ b/src/eval.rs @@ -1,195 +1,244 @@ +use num::bigint::Sign; +use num::{BigInt, FromPrimitive, Integer, Signed, ToPrimitive, Zero}; use std::fmt::{Display, Formatter}; -use Num::{Float, Int}; -#[derive(Clone, Copy, Debug, PartialEq)] +fn int_to_float(i: &BigInt) -> f64 { + match i.sign() { + Sign::Minus => return -int_to_float(&-i), + Sign::NoSign => return 0., + _ => (), + } + let mut exponent = i.bits() - 1; + if exponent > 1023 { + return f64::INFINITY; + } + let mut fraction; + if exponent <= 52 { + fraction = i.to_u64().unwrap() << 52 - exponent; + } else { + let power = BigInt::from(1) << exponent - 52; + let (fraction_big, rem) = i.div_rem(&power); + fraction = fraction_big.to_u64().unwrap(); + let half = power >> 1; + if rem > half || rem == half && fraction % 2 == 1 { + fraction += 1; + if fraction == 1 << 53 { + exponent += 1; + if exponent > 1023 { + return f64::INFINITY; + } + fraction = 1 << 53; + } + } + } + + let u = (exponent + 1023) << 52 | fraction & !(1 << 52); + f64::from_bits(u) +} + +#[derive(Debug, PartialEq)] pub enum Num { - Int(i128), + Int(BigInt), Float(f64), } +fn int(i: BigInt) -> Num { + if i.bits() > 1 << 30 { + return Num::Float(int_to_float(&i)); + } + Num::Int(i) +} + +fn float(f: f64) -> Num { + Num::Float(f) +} + impl Num { - fn as_float(self) -> f64 { + fn as_float(&self) -> f64 { match self { - Int(i) => i as f64, - Float(f) => f, + Num::Int(i) => int_to_float(i), + &Num::Float(f) => f, } } - fn powi(self, p: i128) -> Num { - if p == 0 { - return Int(1); - } - if p == 1 { - return self; + fn square(self) -> Num { + match self { + Num::Int(i) => int(i.pow(2)), + Num::Float(f) => float(f * f), } - if p % 2 == 0 { - return self.mul(self).powi(p / 2); + } + + fn powi(self, p: BigInt) -> Num { + let mut n = self; + let mut p = p; + let mut acc = int(BigInt::from(1)); + loop { + if p.is_odd() { + acc = acc.mul(&n); + } + p >>= 1; + if p.is_zero() { + return acc; + } + n = n.square(); } - self.mul(self).powi(p / 2).mul(self) } pub fn pow(self, other: Num) -> Num { match other { - Int(i) => { - if i >= 0 { + Num::Int(i) => { + if !i.is_negative() { return self.powi(i); } - if let Ok(i_32) = i32::try_from(i) { - return Float(self.as_float().powi(i_32)); + if let Ok(i_32) = i32::try_from(&i) { + return float(self.as_float().powi(i_32)); } - Float(1.).div(self).powi(-i) + float(1.).div(self).powi(-i) } - Float(f) => Float(self.as_float().powf(f)), + Num::Float(f) => float(self.as_float().powf(f)), } } - pub fn mul(self, other: Num) -> Num { - if let (Int(i1), Int(i2)) = (self, other) { - if let Some(p) = i1.checked_mul(i2) { - return Int(p); - } + pub fn mul(self, other: &Num) -> Num { + match (self, other) { + (Num::Int(i1), Num::Int(i2)) => int(i1 * i2), + (n1, n2) => float(n1.as_float() * n2.as_float()), } - Float(self.as_float() * other.as_float()) } pub fn div(self, other: Num) -> Num { - Float(self.as_float() / other.as_float()) + float(self.as_float() / other.as_float()) } pub fn int_div(self, other: Num) -> Num { - if let (Int(i1), Int(i2)) = (self, other) { - if let Some(q) = i1.checked_div(i2) { - return Int(q); + match (self, other) { + (Num::Int(i1), Num::Int(i2)) if !i2.is_zero() => int(i1.div_floor(&i2)), + (n1, n2) => { + let f = (n1.as_float() / n2.as_float()).floor(); + let Some(i) = BigInt::from_f64(f) else { + return float(f); + }; + int(i) } } - Int((self.as_float() / other.as_float()) as i128) } pub fn modulo(self, other: Num) -> Num { - if let (Int(i1), Int(i2)) = (self, other) { - if let Some(r) = i1.checked_rem(i2) { - if i2 > 0 && r < 0 || i2 < 0 && r > 0 { - return Int(r + i2); - } - return Int(r); - } + match (self, other) { + (Num::Int(i1), Num::Int(i2)) if !i2.is_zero() => int(i1.mod_floor(&i2)), + (n1, n2) => float(n1.as_float() % n2.as_float()), } - let n1 = self.as_float(); - let n2 = other.as_float(); - let r = n1 % n2; - if n2 > 0. && r < 0. || n2 < 0. && r > 0. { - return Float(r + n2); - } - Float(r) } pub fn add(self, other: Num) -> Num { - if let (Int(i1), Int(i2)) = (self, other) { - if let Some(s) = i1.checked_add(i2) { - return Int(s); - } + match (self, other) { + (Num::Int(i1), Num::Int(i2)) => int(i1 + i2), + (n1, n2) => float(n1.as_float() + n2.as_float()), } - Float(self.as_float() + other.as_float()) } pub fn sub(self, other: Num) -> Num { - if let (Int(i1), Int(i2)) = (self, other) { - if let Some(d) = i1.checked_sub(i2) { - return Int(d); - } + match (self, other) { + (Num::Int(i1), Num::Int(i2)) => int(i1 - i2), + (n1, n2) => float(n1.as_float() - n2.as_float()), } - Float(self.as_float() - other.as_float()) } pub fn sin(self) -> Num { - Float(self.as_float().sin()) + float(self.as_float().sin()) } pub fn cos(self) -> Num { - Float(self.as_float().cos()) + float(self.as_float().cos()) } pub fn tan(self) -> Num { - Float(self.as_float().tan()) + float(self.as_float().tan()) } pub fn asin(self) -> Num { - Float(self.as_float().asin()) + float(self.as_float().asin()) } pub fn acos(self) -> Num { - Float(self.as_float().acos()) + float(self.as_float().acos()) } pub fn atan(self) -> Num { - Float(self.as_float().atan()) + float(self.as_float().atan()) } pub fn sqrt(self) -> Num { match self { - Int(i) if i >= 0 => { - if i <= 1 { - return Int(i); - } - let mut x0 = i / 2; - let mut x1 = (x0 + i / x0) / 2; - while x1 < x0 { - x0 = x1; - x1 = (x0 + i / x0) / 2; - } - if x0 * x0 == i { - return Int(x0); + Num::Int(i) if !i.is_negative() => { + let s = i.sqrt(); + if &s * &s == i { + return int(s); } + float(int_to_float(&i).sqrt()) } - _ => (), + n => float(n.as_float().sqrt()), } - Float(self.as_float().sqrt()) } pub fn log(self) -> Num { - Float(self.as_float().ln()) + float(self.as_float().ln()) } pub fn log10(self) -> Num { - Float(self.as_float().log10()) + float(self.as_float().log10()) } pub fn log2(self) -> Num { - Float(self.as_float().log2()) + float(self.as_float().log2()) } pub fn floor(self) -> Num { match self { - Int(i) => Int(i), - Float(f) => Int(f.floor() as i128), + Num::Int(i) => int(i), + Num::Float(f) => { + let f = f.floor(); + if let Some(i) = BigInt::from_f64(f) { + int(i) + } else { + float(f) + } + } } } pub fn ceil(self) -> Num { match self { - Int(i) => Int(i), - Float(f) => Int(f.ceil() as i128), + Num::Int(i) => int(i), + Num::Float(f) => { + let f = f.ceil(); + if let Some(i) = BigInt::from_f64(f) { + int(i) + } else { + float(f) + } + } } } pub fn round(self) -> Num { match self { - Int(i) => Int(i), - Float(f) => Int(f.round_ties_even() as i128), - } - } - - pub fn trunc(self) -> Num { - match self { - Int(i) => Int(i), - Float(f) => Int(f as i128), + Num::Int(i) => int(i), + Num::Float(f) => { + let f = f.round_ties_even(); + if let Some(i) = BigInt::from_f64(f) { + int(i) + } else { + float(f) + } + } } } pub fn abs(self) -> Num { match self { - Int(i) => Int(i.abs()), - Float(f) => Float(f.abs()), + Num::Int(i) => int(i.abs()), + Num::Float(f) => float(f.abs()), } } } @@ -197,16 +246,13 @@ impl Num { impl Display for Num { fn fmt(&self, f: &mut Formatter<'_>) -> Result<(), std::fmt::Error> { match self { - Int(n) => write!(f, "{n}"), - Float(n) => { - if !(-1e30..=1e30).contains(n) { - return write!(f, "{n:.7e}"); + Num::Int(n) if n.bits() <= 100 => write!(f, "{n}"), + n => { + let n = n.as_float(); + if (-(1i128 << 100) as f64) < n && n < (1i128 << 100) as f64 { + return write!(f, "{}", format!("{n:.7}").trim_end_matches('0')); } - let mut buf = format!("{n:.7}"); - while let Some(b'0') = buf.as_bytes().get(buf.len() - 1) { - buf.truncate(buf.len() - 1); - } - write!(f, "{buf}") + write!(f, "{n:.7e}") } } } @@ -216,78 +262,146 @@ impl Display for Num { mod tests { use super::*; + fn int(i: i128) -> Num { + super::int(BigInt::from(i)) + } + #[test] - fn pow_positive_int() { - assert_eq!(Int(2).pow(Int(16)), Int(65536)); + fn int_to_float_small() { + for n in -(1 << 16)..1 << 16 { + assert_eq!( + int_to_float(&BigInt::from(n)), + n as f64, + "int_to_float({}) is not equal to {}", + n, + n as f64 + ); + } } #[test] - fn pow_negative_int() { - assert_eq!(Int(2).pow(Int(-3)), Float(0.125)); + fn int_to_float_medium() { + let n = BigInt::from(2).pow(1023) + BigInt::from(2).pow(1022); + assert_eq!(int_to_float(&n), n.to_f64().unwrap()); } #[test] - fn pow_float() { - assert_eq!(Int(4).pow(Float(0.5)), Float(2.)); + fn int_to_float_round_up() { + assert_eq!( + int_to_float(&BigInt::from(0x7f_ffff_ffff_ffffu64)), + 0x80_0000_0000_0000u64 as f64 + ); } #[test] - fn pow_overflow() { - assert_eq!(Int(2).pow(Int(1 << 126)), Float(f64::INFINITY)); + fn int_to_float_round_up_easy() { + assert_eq!( + int_to_float(&BigInt::from(0x40_0000_0000_0003u64)), + 0x40_0000_0000_0004u64 as f64 + ); } #[test] - fn pow_underflow() { - assert_eq!(Int(2).pow(Int(-(1 << 126))), Float(0.)); + fn int_to_float_round_down() { + assert_eq!( + int_to_float(&BigInt::from(0x40_0000_0000_0001u64)), + 0x40_0000_0000_0000u64 as f64 + ); } #[test] - fn modulo_pos() { - assert_eq!(Int(5).modulo(Int(3)), Int(2)); + fn int_to_float_round_to_even_down() { + assert_eq!( + int_to_float(&BigInt::from(0x40_0000_0000_0002u64)), + 0x40_0000_0000_0000u64 as f64 + ); } #[test] - fn modulo_pos_neg() { - assert_eq!(Int(5).modulo(Int(-3)), Int(-1)); + fn int_to_float_round_to_even_up() { + assert_eq!( + int_to_float(&BigInt::from(0x40_0000_0000_0006u64)), + 0x40_0000_0000_0008u64 as f64 + ); } #[test] - fn modulo_neg_pos() { - assert_eq!(Int(-5).modulo(Int(3)), Int(1)); + fn int_to_float_large() { + assert_eq!( + int_to_float( + &((BigInt::from(0x1f_ffff_ffff_ffffu64) << 971) + (BigInt::from(1) << 970) - 1) + ), + f64::MAX + ); } #[test] - fn modulo_neg() { - assert_eq!(Int(-5).modulo(Int(-3)), Int(-2)); + fn int_to_float_overflow() { + assert_eq!( + int_to_float( + &((BigInt::from(0x1f_ffff_ffff_ffffu64) << 971) + (BigInt::from(1) << 970)) + ), + f64::INFINITY + ); + } + + #[test] + fn pow_positive_int() { + assert_eq!(int(2).pow(int(16)), int(65536)); + } + + #[test] + fn pow_negative_int() { + assert_eq!(int(2).pow(int(-3)), float(0.125)); + } + + #[test] + fn pow_float() { + assert_eq!(int(4).pow(float(0.5)), float(2.)); + } + + #[test] + fn pow_overflow() { + assert_eq!(int(2).pow(int(1 << 126)), float(f64::INFINITY)); + } + + #[test] + fn pow_underflow() { + assert_eq!(int(2).pow(int(-(1 << 126))), float(0.)); + } + + #[test] + fn modulo_pos() { + assert_eq!(int(5).modulo(int(3)), int(2)); } #[test] - fn modulo_float_pos() { - assert_eq!(Float(5.).modulo(Int(3)), Float(2.)); + fn modulo_pos_neg() { + assert_eq!(int(5).modulo(int(-3)), int(-1)); } #[test] - fn modulo_float_pos_neg() { - assert_eq!(Float(5.).modulo(Int(-3)), Float(-1.)); + fn modulo_neg_pos() { + assert_eq!(int(-5).modulo(int(3)), int(1)); } #[test] - fn modulo_float_neg_pos() { - assert_eq!(Float(-5.).modulo(Int(3)), Float(1.)); + fn modulo_neg() { + assert_eq!(int(-5).modulo(int(-3)), int(-2)); } #[test] - fn modulo_float_neg() { - assert_eq!(Float(-5.).modulo(Int(-3)), Float(-2.)); + fn modulo_float() { + assert_eq!(float(5.).modulo(int(3)), float(2.)); } #[test] fn sqrt() { for n in 0..65536 { assert_eq!( - Int(n * n).sqrt(), - Int(n), - "Int({}).sqrt() is not equal to {}", + int(n * n).sqrt(), + int(n), + "int({}).sqrt() is not equal to {}", n * n, n ); @@ -296,6 +410,6 @@ mod tests { #[test] fn sqrt_big() { - assert_eq!(Int(1 << 126).sqrt(), Int(1 << 63)); + assert_eq!(int(1 << 126).sqrt(), int(1 << 63)); } } |
