#![allow(clippy::cast_possible_truncation)]
#[inline(always)]
pub fn sat_i8(v: i32) -> i8 {
if v > i8::MAX as i32 {
i8::MAX
} else if v < i8::MIN as i32 {
i8::MIN
} else {
v as i8
}
}
pub fn quantize_multiplier(m_real: f32) -> (i32, i32) {
assert!(
m_real.is_finite() && m_real > 0.0,
"quantize_multiplier: m_real must be finite and positive (got {m_real})"
);
assert!(
m_real < 1.0,
"quantize_multiplier: m_real ≥ 1 unsupported (got {m_real}); \
FPGA epilogue assumes shrinking requant"
);
let bits = m_real.to_bits();
let frac = (bits & 0x007F_FFFF) as i32;
let exp = ((bits >> 23) & 0xFF) as i32;
let exp_norm = exp - 126;
let m0 = (1i32 << 30) + (frac << 7);
let shift = -exp_norm;
let shift = shift.clamp(0, 31);
(m0, shift)
}
#[inline]
pub fn srdhm(a: i32, b: i32) -> i32 {
if a == i32::MIN && b == i32::MIN {
return i32::MAX;
}
let ab = (a as i64) * (b as i64);
let nudge: i64 = if ab >= 0 { 1 << 30 } else { 1 - (1 << 30) };
((ab + nudge) / (1i64 << 31)) as i32
}
#[inline]
pub fn rdpot(x: i32, shift: i32) -> i32 {
debug_assert!((0..=31).contains(&shift));
if shift == 0 {
return x;
}
let mask = (1i32 << shift) - 1;
let remainder = x & mask;
let threshold = (mask >> 1) + i32::from(x < 0);
(x >> shift) + i32::from(remainder > threshold)
}
#[inline]
pub fn requantize_q31(acc: i32, m0: i32, shift: i32, out_zp: i32) -> i8 {
let prod = srdhm(acc, m0);
let shifted = rdpot(prod, shift);
sat_i8(shifted + out_zp)
}
#[inline]
pub fn q31_to_q15(m0_q31: i32, shift: i32) -> (i16, i32) {
let m0_q15 = ((m0_q31 as i64 + (1 << 15)) >> 16) as i32;
let (m0_q15, shift) = if m0_q15 >= (1 << 15) {
(m0_q15 / 2, shift - 1)
} else {
(m0_q15, shift)
};
let m0_q15 = m0_q15.clamp(i16::MIN as i32, i16::MAX as i32) as i16;
(m0_q15, shift)
}
#[inline]
pub fn srdhm_q15(a: i32, b: i16) -> i32 {
let ab = (a as i64) * (b as i64);
let nudge: i64 = if ab >= 0 { 1 << 14 } else { 1 - (1 << 14) };
((ab + nudge) / (1i64 << 15)) as i32
}
#[inline]
pub fn requantize_q15(acc: i32, m0: i16, shift: i32, out_zp: i32) -> i8 {
let prod = srdhm_q15(acc, m0);
let shifted = rdpot(prod, shift);
sat_i8(shifted + out_zp)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn quantize_half() {
let (m0, sh) = quantize_multiplier(0.5);
assert_eq!(m0, 1i32 << 30);
assert_eq!(sh, 0);
}
#[test]
fn quantize_quarter() {
let (m0, sh) = quantize_multiplier(0.25);
assert_eq!(m0, 1i32 << 30);
assert_eq!(sh, 1);
}
#[test]
fn m0_in_q031_range() {
for m in [1e-3, 1e-2, 1e-1, 0.3, 0.49, 0.99, 1e-6, 1e-9] {
let (m0, _) = quantize_multiplier(m);
assert!(
m0 >= 1 << 30 && m0 < i32::MAX,
"m={m}: M0={m0} out of [2^30, 2^31)"
);
}
}
#[test]
fn requant_close_to_real_multiplier() {
let mults = [1e-3_f32, 7.87e-3, 3e-2, 0.1, 0.5, 0.99];
for &m in &mults {
let (m0, shift) = quantize_multiplier(m);
for acc in [-100_000_i32, -1, 0, 1, 100, 12_345, 100_000] {
let got = requantize_q31(acc, m0, shift, 0) as i32;
let want = (acc as f64 * m as f64).round() as i32;
let want_clamped = want.clamp(i8::MIN as i32, i8::MAX as i32);
let diff = (got - want_clamped).abs();
assert!(
diff <= 1,
"m={m} acc={acc}: got {got}, want {want_clamped} (diff {diff})"
);
}
}
}
#[test]
fn srdhm_saturates_at_min_min() {
assert_eq!(srdhm(i32::MIN, i32::MIN), i32::MAX);
}
#[test]
fn srdhm_half_is_divide_by_two_no_ties() {
let half = 1i32 << 30;
assert_eq!(srdhm(0, half), 0);
assert_eq!(srdhm(4, half), 2);
assert_eq!(srdhm(-4, half), -2);
assert_eq!(srdhm(100, half), 50);
assert_eq!(srdhm(-100, half), -50);
}
#[test]
fn rdpot_rounds_half_away_from_zero() {
assert_eq!(rdpot(5, 1), 3);
assert_eq!(rdpot(-5, 1), -3);
assert_eq!(rdpot(4, 1), 2);
assert_eq!(rdpot(12345, 0), 12345);
}
#[test]
fn q15_within_one_ulp_of_q31() {
let mults = [1e-3_f32, 7.87e-3, 3e-2, 0.1, 0.5, 0.9];
for &m in &mults {
let (m0_q31, sh_q31) = quantize_multiplier(m);
let (m0_q15, sh_q15) = q31_to_q15(m0_q31, sh_q31);
for acc in [-10_000_i32, -1, 0, 1, 100, 9_999, 50_000] {
let want = requantize_q31(acc, m0_q31, sh_q31, 0) as i32;
let got = requantize_q15(acc, m0_q15, sh_q15, 0) as i32;
assert!(
(want - got).abs() <= 1,
"m={m} acc={acc}: q31={want}, q15={got}"
);
}
}
}
}