use rug::Integer;
use crate::rfloat::RFloatContext;
use crate::util::bitmask;
use crate::{Real, RoundingContext, RoundingMode, Split};
use super::{Posit, PositVal};
#[derive(Clone, Debug)]
pub struct PositContext {
es: usize,
nbits: usize,
}
impl PositContext {
pub const ES_MAX: usize = 32;
pub const PAD_MIN: usize = 3;
pub fn new(es: usize, nbits: usize) -> Self {
assert!(
es <= Self::ES_MAX,
"exponent width needs to be at most {} bits, given {} bits",
Self::ES_MAX,
es
);
assert!(
nbits >= es + Self::PAD_MIN,
"total bitwidth needs to be at least {} bits, given {} bits",
es + Self::PAD_MIN,
nbits
);
Self { es, nbits }
}
pub fn es(&self) -> usize {
self.es
}
pub fn nbits(&self) -> usize {
self.nbits
}
pub fn max_p(&self) -> usize {
self.nbits - self.es - 3
}
pub fn useed(&self) -> isize {
(1_usize << (1 << self.es)) as isize
}
pub fn rscale(&self) -> isize {
(1 << self.es) as isize
}
pub fn rmax(&self) -> isize {
let max_r = (self.nbits - 1) as isize;
max_r - 1
}
pub fn emax(&self) -> isize {
self.rscale() * self.rmax()
}
pub fn expmax(&self) -> isize {
self.emax()
}
pub fn emin(&self) -> isize {
self.rscale() * -self.rmax()
}
pub fn expmin(&self) -> isize {
self.emin() }
pub fn maxval(&self, sign: bool) -> Posit {
Posit {
num: PositVal::NonZero(sign, self.rmax(), 0, Integer::from(1)),
ctx: self.clone(),
}
}
pub fn minval(&self, sign: bool) -> Posit {
Posit {
num: PositVal::NonZero(sign, -self.rmax(), 0, Integer::from(1)),
ctx: self.clone(),
}
}
pub fn zero(&self) -> Posit {
Posit {
num: PositVal::Zero,
ctx: self.clone(),
}
}
pub fn nar(&self) -> Posit {
Posit {
num: PositVal::Nar,
ctx: self.clone(),
}
}
pub fn bits_to_number(&self, b: Integer) -> Posit {
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 ns = b & bitmask(self.nbits - 1);
if ns == 0 {
Posit {
num: if s { PositVal::Nar } else { PositVal::Zero },
ctx: self.clone(),
}
} else {
let r0 = ns.get_bit((self.nbits - 2) as u32);
let mut r0_pos = self.nbits - 2;
while r0_pos > 0 && ns.get_bit((r0_pos - 1) as u32) == r0 {
r0_pos -= 1;
}
if r0_pos == 0 {
Posit {
num: PositVal::NonZero(s, self.rmax(), 0, Integer::from(1)),
ctx: self.clone(),
}
} else {
let embits = r0_pos - 1;
let rbits = self.nbits - embits - 1;
let (ebits, mbits) = if embits <= self.es {
(embits, 0)
} else {
(self.es, embits - self.es)
};
let efield = (ns.clone() >> mbits) & bitmask(ebits);
let mfield = ns & bitmask(mbits);
let kbits = rbits - 1;
let regime = if r0 {
kbits as isize - 1
} else {
-(kbits as isize)
};
let e = if ebits < self.es {
efield.to_isize().unwrap() << (self.es - ebits)
} else {
efield.to_isize().unwrap()
};
let c = mfield | (1 << mbits);
Posit {
num: PositVal::NonZero(s, regime, e - mbits as isize, c),
ctx: self.clone(),
}
}
}
}
}
impl PositContext {
fn round_params<T: Real>(&self, num: &T) -> (isize, usize) {
assert!(
!num.is_nar() && !num.is_zero(),
"must be a finite, non-zero {:?}",
num
);
let useed = self.useed();
let r = num.e().unwrap() / useed;
let kbits = if r < 0 { -r } else { r + 1 } as usize;
let embits = self.nbits - (kbits + 2);
let mbits = if embits <= self.es {
0
} else {
embits - self.es
};
(useed, mbits)
}
fn round_finite(&self, split: Split, useed: isize) -> Posit {
let s = split.sign().unwrap();
let rounded = RFloatContext::round_finalize(split, RoundingMode::NearestTiesToEven);
let e = rounded.e().unwrap();
let r = e / useed;
let e = e % useed;
let c = rounded.c().unwrap();
let exp = (e + 1) - (c.significant_bits() as isize);
Posit {
num: PositVal::NonZero(s, r, exp, c),
ctx: self.clone(),
}
}
}
impl RoundingContext for PositContext {
type Format = Posit;
fn round<T: Real>(&self, val: &T) -> Self::Format {
if val.is_nar() {
self.nar()
} else if val.is_zero() {
self.zero()
} else {
let s = val.sign().unwrap();
let e = val.e().unwrap();
if e >= self.emax() {
self.maxval(s)
} else if e <= self.emin() {
self.minval(s)
} else {
let (useed, mbits) = self.round_params(val);
let (p, n) = RFloatContext::new().with_max_p(mbits + 1).round_params(val);
let split = Split::new(val, p, n);
self.round_finite(split, useed)
}
}
}
}