use eunomia::{Bf16, CastFrom, FloatElement, NumericElement, F16, F32, F64};
macro_rules! float_element_contract {
($name:ident, $t:ty) => {
#[test]
fn $name() {
let f = |v: f32| <$t as FloatElement>::from_f32(v);
let g = |x: $t| FloatElement::to_f32(x);
assert_eq!(g(<$t as NumericElement>::ZERO), 0.0, "ZERO");
assert_eq!(g(<$t as NumericElement>::ONE), 1.0, "ONE");
assert_eq!(
g(<$t as NumericElement>::INFINITY),
f32::INFINITY,
"INFINITY"
);
assert_eq!(
<$t as NumericElement>::BYTE_WIDTH,
core::mem::size_of::<$t>(),
"BYTE_WIDTH == size_of"
);
assert!(
NumericElement::is_nan(<$t as NumericElement>::NAN),
"is_nan(NAN)"
);
assert!(
!NumericElement::is_finite(<$t as NumericElement>::NAN),
"!is_finite(NAN)"
);
assert!(
!NumericElement::is_finite(<$t as NumericElement>::INFINITY),
"!is_finite(INFINITY)"
);
assert!(
NumericElement::is_finite(<$t as NumericElement>::ONE),
"is_finite(ONE)"
);
assert!(
!NumericElement::is_nan(<$t as NumericElement>::ONE),
"!is_nan(ONE)"
);
let a = f(1.5);
assert_eq!(a + <$t as NumericElement>::ZERO, a, "a + 0 == a");
assert_eq!(a * <$t as NumericElement>::ONE, a, "a * 1 == a");
assert_eq!(g(f(1.5) + f(2.5)), 4.0, "1.5 + 2.5");
assert_eq!(g(f(4.0) - f(2.5)), 1.5, "4.0 - 2.5");
assert_eq!(g(f(2.0) * f(3.0)), 6.0, "2.0 * 3.0");
assert_eq!(g(f(6.0) / f(2.0)), 3.0, "6.0 / 2.0");
assert_eq!(g(NumericElement::abs(f(-2.5))), 2.5, "abs(-2.5)");
assert_eq!(g(NumericElement::sqrt(f(4.0))), 2.0, "sqrt(4)");
assert_eq!(g(FloatElement::signum(f(-2.5))), -1.0, "signum(-2.5)");
assert_eq!(g(FloatElement::signum(f(2.5))), 1.0, "signum(2.5)");
assert_eq!(
g(NumericElement::scalar_fmadd(f(2.0), f(3.0), f(1.0))),
7.0,
"fmadd 2*3+1"
);
assert_eq!(g(NumericElement::min_scalar(f(1.5), f(2.5))), 1.5, "min");
assert_eq!(g(NumericElement::max_scalar(f(1.5), f(2.5))), 2.5, "max");
assert_eq!(g(<$t as FloatElement>::from_f64(1.5_f64)), 1.5, "from_f64");
assert_eq!(g(FloatElement::powi(f(2.0), 3)), 8.0, "2^3");
assert_eq!(g(FloatElement::powi(f(2.0), 0)), 1.0, "2^0");
assert_eq!(g(FloatElement::powi(f(2.0), -1)), 0.5, "2^-1");
assert_eq!(
g(FloatElement::exp(<$t as NumericElement>::ZERO)),
1.0,
"exp(0)"
);
assert_eq!(
g(FloatElement::ln(<$t as NumericElement>::ONE)),
0.0,
"ln(1)"
);
assert_eq!(
g(FloatElement::sin(<$t as NumericElement>::ZERO)),
0.0,
"sin(0)"
);
assert_eq!(
g(FloatElement::cos(<$t as NumericElement>::ZERO)),
1.0,
"cos(0)"
);
assert_eq!(
g(<$t as CastFrom<i32>>::cast_from(5_i32)),
5.0,
"cast_from(5)"
);
}
};
}
float_element_contract!(f16_element_contract, F16);
float_element_contract!(bf16_element_contract, Bf16);
float_element_contract!(f32_wrapper_element_contract, F32);
float_element_contract!(f64_wrapper_element_contract, F64);
#[test]
fn reduced_precision_rounds_to_the_native_grid() {
let f16 = FloatElement::to_f32(<F16 as FloatElement>::from_f32(0.1));
let bf16 = FloatElement::to_f32(<Bf16 as FloatElement>::from_f32(0.1));
assert!(
(f16 - 0.1).abs() <= 0.1 * 2.0_f32.powi(-10),
"F16 |err| ≤ rel 2^-10"
);
assert!(
(bf16 - 0.1).abs() <= 0.1 * 2.0_f32.powi(-7),
"Bf16 |err| ≤ rel 2^-7"
);
assert!(
(f16 - 0.1).abs() < (bf16 - 0.1).abs(),
"F16 grid finer than Bf16 at 0.1"
);
assert_ne!(f16, bf16, "F16 and Bf16 quantize 0.1 to distinct values");
}