use core::mem::size_of;
use crate::{
constants::{EARLY_EXIT_BIT_WIDTH, SAMPLES_COUNT},
float::AlpFloat,
};
const PRE_CHECK_LEN_MUL: usize = 6;
const PRE_CHECK_LEN_FAC: usize = 4;
const PRE_CHECK_MAX_EXC_FAC: usize = 2;
const DIV_EARLY_CHECK_LEN: usize = 6;
const DIV_EARLY_ABORT_EXC: usize = 3;
const FAC_PENALTY_MULT: usize = 2;
const LOW_COST_THRESHOLD_PER_VAL: usize = 3;
const HIGH_EXP_DIV_THRESHOLD: u8 = 14;
const EXP_PRIORITY: [u8; 19] = [
2, 1, 3, 0, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18,
];
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct BestParams {
pub exp: u8,
pub fac: u8,
pub use_div: bool,
}
#[inline]
pub(crate) fn find_identical_base<F: AlpFloat>(val: F) -> Option<(u8, F::Int)> {
const FAC_INT: i64 = 1;
(0..=F::MAX_EXPONENT).find_map(|exp| {
let frac_exp = F::frac_exp(exp);
let exp_factor = F::exp_factor(exp, 0);
F::try_encode_fast(val, exp_factor, FAC_INT, frac_exp).map(|base| (exp, base))
})
}
pub(crate) fn find_best_params<F: AlpFloat>(samples: &[F]) -> BestParams {
if samples.is_empty() {
return BestParams {
exp: 0,
fac: 0,
use_div: false,
};
}
let mut valid_samples: [F; SAMPLES_COUNT] = [F::ZERO; SAMPLES_COUNT];
let mut sample_len = 0;
for &val in samples
.iter()
.filter(|v| !v.is_impossible())
.take(SAMPLES_COUNT)
{
valid_samples[sample_len] = val;
sample_len += 1;
}
let active_samples = &valid_samples[..sample_len];
if sample_len == 0 {
return BestParams {
exp: 0,
fac: 0,
use_div: false,
};
}
let mut best_cost = size_of::<F>() * 8 * sample_len;
let mut best_exceptions = sample_len;
let mut best_params = BestParams {
exp: 0,
fac: 0,
use_div: false,
};
let mut any_decimal = false;
for &exp in &EXP_PRIORITY {
if exp > F::MAX_EXPONENT {
continue;
}
let frac_exp = F::frac_exp(exp);
let exp_factor = F::exp_factor(exp, 0);
const FAC_INT: i64 = 1;
let pre_n = active_samples.len().min(PRE_CHECK_LEN_MUL);
let mut pre_exc = 0;
let mut min_val = F::MAX_INT;
let mut max_val = F::MIN_INT;
for &val in &active_samples[..pre_n] {
if let Some(enc) = F::try_encode_fast(val, exp_factor, FAC_INT, frac_exp) {
min_val = min_val.min(enc);
max_val = max_val.max(enc);
} else {
pre_exc += 1;
}
}
let mul_feasible = pre_exc < pre_n;
let mut exceptions = pre_exc;
if mul_feasible {
for &val in &active_samples[pre_n..] {
if let Some(enc) = F::try_encode_fast(val, exp_factor, FAC_INT, frac_exp) {
min_val = min_val.min(enc);
max_val = max_val.max(enc);
} else {
exceptions += 1;
if exceptions * F::EXCEPTION_PENALTY >= best_cost {
break;
}
}
}
if exceptions != sample_len {
any_decimal = true;
if exceptions * F::EXCEPTION_PENALTY < best_cost {
let max_offset = if min_val <= max_val {
F::calc_range(min_val, max_val)
} else {
0
};
let bit_width = F::bits_needed(max_offset) as usize;
let total_cost = bit_width * sample_len + exceptions * F::EXCEPTION_PENALTY;
if total_cost < best_cost {
best_cost = total_cost;
best_exceptions = exceptions;
best_params = BestParams {
exp,
fac: 0,
use_div: false,
};
if total_cost == 0 || (exceptions == 0 && bit_width <= EARLY_EXIT_BIT_WIDTH) {
return best_params;
}
}
}
}
}
if exp > 0 && (!mul_feasible || exceptions > 0 || exp >= HIGH_EXP_DIV_THRESHOLD) {
let mut div_exceptions = 0usize;
let mut div_min = F::MAX_INT;
let mut div_max = F::MIN_INT;
for (idx, &val) in active_samples.iter().enumerate() {
if let Some(enc) = F::try_encode_div(val, exp_factor) {
div_min = div_min.min(enc);
div_max = div_max.max(enc);
} else {
div_exceptions += 1;
if div_exceptions >= DIV_EARLY_ABORT_EXC && idx < DIV_EARLY_CHECK_LEN {
div_exceptions = sample_len;
break;
}
if div_exceptions * F::EXCEPTION_PENALTY >= best_cost {
break;
}
}
}
if div_exceptions != sample_len {
any_decimal = true;
if div_exceptions * F::EXCEPTION_PENALTY < best_cost {
let max_offset = if div_min <= div_max {
F::calc_range(div_min, div_max)
} else {
0
};
let bit_width = F::bits_needed(max_offset) as usize;
let total_cost = bit_width * sample_len + div_exceptions * F::EXCEPTION_PENALTY;
if total_cost < best_cost {
best_cost = total_cost;
best_exceptions = div_exceptions;
best_params = BestParams {
exp,
fac: 0,
use_div: true,
};
if total_cost == 0 || (div_exceptions == 0 && bit_width <= EARLY_EXIT_BIT_WIDTH) {
return best_params;
}
}
}
}
}
}
if !any_decimal || best_cost <= sample_len * LOW_COST_THRESHOLD_PER_VAL {
return best_params;
}
let max_search_exp = if best_exceptions == 0 {
best_params.exp.saturating_sub(1)
} else {
F::MAX_EXPONENT
};
if max_search_exp == 0 {
return best_params;
}
for exp in 1..=max_search_exp {
let max_fac = exp.min(F::MAX_FAC);
let frac_exp = F::frac_exp(exp);
for fac in 1..=max_fac {
let exp_factor = F::exp_factor(exp, fac);
let fac_int = F::fac_int(fac);
let pre_n = active_samples.len().min(PRE_CHECK_LEN_FAC);
let mut pre_exc = 0;
let mut min_val = F::MAX_INT;
let mut max_val = F::MIN_INT;
for &val in &active_samples[..pre_n] {
if let Some(enc) = F::try_encode_fast(val, exp_factor, fac_int, frac_exp) {
min_val = min_val.min(enc);
max_val = max_val.max(enc);
} else {
pre_exc += 1;
}
}
let fac_penalty = sample_len * FAC_PENALTY_MULT;
if pre_exc >= PRE_CHECK_MAX_EXC_FAC
|| pre_exc * F::EXCEPTION_PENALTY + fac_penalty >= best_cost
{
continue;
}
let mut exceptions = pre_exc;
for &val in &active_samples[pre_n..] {
if let Some(enc) = F::try_encode_fast(val, exp_factor, fac_int, frac_exp) {
min_val = min_val.min(enc);
max_val = max_val.max(enc);
} else {
exceptions += 1;
if exceptions * F::EXCEPTION_PENALTY + fac_penalty >= best_cost {
break;
}
}
}
if exceptions != sample_len && exceptions * F::EXCEPTION_PENALTY + fac_penalty < best_cost {
let max_offset = if min_val <= max_val {
F::calc_range(min_val, max_val)
} else {
0
};
let bit_width = F::bits_needed(max_offset) as usize;
let total_cost = bit_width * sample_len + exceptions * F::EXCEPTION_PENALTY + fac_penalty;
if total_cost < best_cost {
best_cost = total_cost;
best_params = BestParams {
exp,
fac,
use_div: false,
};
if total_cost == 0 || (exceptions == 0 && bit_width <= EARLY_EXIT_BIT_WIDTH) {
return best_params;
}
}
}
}
}
best_params
}