use crate::{
eval::{
instructions::{BinaryOp, TernaryOp, UnaryOp, UnaryParamOp},
ops,
},
interval::Ival,
mpfr::mpfr_get_exp,
};
use rug::Float;
pub fn get_slack(iteration: usize, slack_unit: i64) -> i64 {
if iteration == 0 || slack_unit <= 0 {
0
} else {
let shift = iteration.saturating_sub(1) as u32;
slack_unit.checked_shl(shift).unwrap_or(i64::MAX)
}
}
pub fn slack_bits(iteration: usize, slack_unit: i64) -> u32 {
clamp_to_bits(get_slack(iteration, slack_unit))
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct AmplBounds {
pub upper: i64,
pub lower: i64,
}
impl AmplBounds {
pub const fn new(upper: i64, lower: i64) -> Self {
Self { upper, lower }
}
pub const fn zero() -> Self {
Self::new(0, 0)
}
}
impl Default for AmplBounds {
fn default() -> Self {
AmplBounds::zero()
}
}
#[derive(Clone, Copy, Debug)]
pub struct TrickContext {
pub iteration: usize,
pub lower_bound_early_stopping: bool,
pub bumps_enabled: bool,
pub slack_unit: i64,
}
impl TrickContext {
pub fn new(
iteration: usize,
lower_bound_early_stopping: bool,
bumps_enabled: bool,
slack_unit: i64,
) -> Self {
Self {
iteration,
lower_bound_early_stopping,
bumps_enabled,
slack_unit,
}
}
pub fn bounds_for_unary(&self, op: UnaryOp, output: &Ival, input: &Ival) -> AmplBounds {
ops::bounds_for_unary(self, op, output, input)
}
pub fn bounds_for_binary(
&self,
op: BinaryOp,
output: &Ival,
lhs: &Ival,
rhs: &Ival,
) -> (AmplBounds, AmplBounds) {
ops::bounds_for_binary(self, op, output, lhs, rhs)
}
pub fn bounds_for_ternary(
&self,
op: TernaryOp,
output: &Ival,
arg1: &Ival,
arg2: &Ival,
arg3: &Ival,
) -> (AmplBounds, AmplBounds, AmplBounds) {
ops::bounds_for_ternary(self, op, output, arg1, arg2, arg3)
}
pub fn bounds_for_unary_param(
&self,
op: UnaryParamOp,
param: u64,
output: &Ival,
input: &Ival,
) -> AmplBounds {
ops::bounds_for_unary_param(self, op, param, output, input)
}
pub fn logspan(&self, value: &Ival) -> i64 {
if !self.bumps_enabled {
return 0;
}
let lo = value.lo.as_float();
let hi = value.hi.as_float();
if lo.is_zero() || hi.is_zero() || lo.is_infinite() || hi.is_infinite() {
return get_slack(self.iteration, self.slack_unit);
}
let lo_exp = exponent(lo);
let hi_exp = exponent(hi);
(lo_exp - hi_exp).abs() + 1
}
pub fn maxlog(&self, value: &Ival, less_slack: bool) -> i64 {
let (lo, hi) = (value.lo.as_float(), value.hi.as_float());
let slack = get_slack(self.iter_for(less_slack), self.slack_unit);
let (lo_inf, hi_inf) = (lo.is_infinite(), hi.is_infinite());
if lo_inf && hi_inf {
slack
} else if hi_inf {
exponent(lo).max(0) + slack
} else if lo_inf {
exponent(hi).max(0) + slack
} else {
exponent(lo).max(exponent(hi)) + 1
}
}
pub fn minlog(&self, value: &Ival, less_slack: bool) -> i64 {
let (lo, hi) = (value.lo.as_float(), value.hi.as_float());
let slack = get_slack(self.iter_for(less_slack), self.slack_unit);
let (lo_zero, hi_zero) = (lo.is_zero(), hi.is_zero());
let (lo_inf, hi_inf) = (lo.is_infinite(), hi.is_infinite());
if lo_zero && hi_zero {
slack
} else if lo_zero {
if hi_inf {
-slack
} else {
exponent(hi).min(0) - slack
}
} else if hi_zero {
if lo_inf {
-slack
} else {
exponent(lo).min(0) - slack
}
} else if crosses_zero(value) {
if hi_inf && lo_inf {
-slack
} else if hi_inf {
exponent(lo).min(0) - slack
} else if lo_inf {
exponent(hi).min(0) - slack
} else {
exponent(lo).min(exponent(hi)).min(0) - slack
}
} else if lo_inf {
exponent(hi)
} else if hi_inf {
exponent(lo)
} else {
exponent(lo).min(exponent(hi))
}
}
fn iter_for(&self, less_slack: bool) -> usize {
if less_slack {
self.iteration.saturating_sub(1)
} else {
self.iteration
}
}
}
pub fn exponent(value: &Float) -> i64 {
if value.is_zero() {
i64::MIN / 4
} else {
mpfr_get_exp(value)
}
}
pub fn crosses_zero(value: &Ival) -> bool {
let lo = value.lo.as_float();
let hi = value.hi.as_float();
let lo_sign = if lo.is_zero() {
0
} else if lo.is_sign_positive() {
1
} else {
-1
};
let hi_sign = if hi.is_zero() {
0
} else if hi.is_sign_positive() {
1
} else {
-1
};
lo_sign != hi_sign
}
pub fn clamp_to_bits(value: i64) -> u32 {
if value <= 0 {
0
} else if value >= u32::MAX as i64 {
u32::MAX
} else {
value as u32
}
}