use cubecl::prelude::*;
use cubecl_core as cubecl;
const SIGN: u32 = 0x8;
const MANTISSA: u32 = 0x1;
const NIBBLE: u32 = 0xF;
const MAGNITUDE: u32 = 0x7;
const MAGNITUDE_SHIFT: u32 = 22;
const EXPONENT_BIAS: u32 = 126 << 23;
const SIGN_SHIFT: u32 = 28;
#[cube]
pub fn e2m1_bits_to_float<F: Numeric, N: Size>(code: Vector<u32, N>) -> Vector<F, N> {
let magnitude = code & Vector::new(MAGNITUDE);
let normal = (magnitude << Vector::new(MAGNITUDE_SHIFT)) + Vector::new(EXPONENT_BIAS);
let mantissa = code & Vector::new(MANTISSA);
let subnormal = select_many(
mantissa.equal(&Vector::new(MANTISSA)),
Vector::new(EXPONENT_BIAS),
Vector::new(0u32),
);
let bits = select_many(
magnitude.greater_than(&Vector::new(MANTISSA)),
normal,
subnormal,
);
let sign = (code & Vector::new(SIGN)) << Vector::new(SIGN_SHIFT);
Vector::<F, N>::cast_from(Vector::<f32, N>::reinterpret(bits | sign))
}
#[cube]
pub fn e2m1_packed_bits_to_float<F: Numeric, N: Size>(word: u32) -> Vector<F, N> {
let mut codes = Vector::<u32, N>::empty();
#[unroll]
for lane in 0..N::value() {
codes.insert(lane, (word >> (4 * lane as u32)) & NIBBLE);
}
e2m1_bits_to_float::<F, N>(codes)
}
#[cube]
pub fn float_to_e2m1_bits<F: Numeric, N: Size>(value: Vector<F, N>) -> Vector<u32, N> {
let value = Vector::<f32, N>::cast_from(value);
let sign_bit = Vector::new(0x8000_0000u32);
let negative = (Vector::<u32, N>::reinterpret(value) & sign_bit).equal(&sign_bit);
let magnitude = select_many(negative, -value, value);
let mut code = cleared::<N>(magnitude.greater_than(&Vector::new(0.25f32)));
code += cleared::<N>(magnitude.greater_equal(&Vector::new(0.75f32)));
code += cleared::<N>(magnitude.greater_than(&Vector::new(1.25f32)));
code += cleared::<N>(magnitude.greater_equal(&Vector::new(1.75f32)));
code += cleared::<N>(magnitude.greater_than(&Vector::new(2.5f32)));
code += cleared::<N>(magnitude.greater_equal(&Vector::new(3.5f32)));
code += cleared::<N>(magnitude.greater_than(&Vector::new(5.0f32)));
code | (cleared::<N>(negative) * Vector::new(SIGN))
}
#[cube]
fn cleared<N: Size>(above: Vector<bool, N>) -> Vector<u32, N> {
select_many(above, Vector::new(1u32), Vector::new(0u32))
}
#[cfg(test)]
mod tests {
use cubecl_common::e2m1;
#[test]
fn the_midpoints_round_to_even() {
for (midpoint, expected) in [
(0.25f32, 0.0f32),
(0.75, 1.0),
(1.25, 1.0),
(1.75, 2.0),
(2.5, 2.0),
(3.5, 4.0),
(5.0, 4.0),
] {
let landed = e2m1::from_f32(midpoint).to_f32();
assert_eq!(landed, expected, "{midpoint} rounded to {landed}");
}
}
#[test]
fn magnitudes_past_the_maximum_saturate() {
for value in [6.0f32, 6.1, 100.0, f32::MAX, f32::INFINITY] {
assert_eq!(e2m1::from_f32(value).to_f32(), 6.0);
assert_eq!(e2m1::from_f32(-value).to_f32(), -6.0);
}
}
#[test]
fn the_sign_bit_survives_a_round_trip() {
for code in 8..16u8 {
let value = e2m1::from_bits(code).to_f32();
assert!(value.is_sign_negative(), "code {code} decoded as {value}");
assert_eq!(e2m1::from_f32(value).to_bits(), code);
}
}
}