use cubecl_core as cubecl;
use cubecl_core::prelude::*;
use cubecl_runtime::runtime::Runtime;
use cubecl_common::e2m1;
use crate::quant::fp4::{e2m1_bits_to_float, e2m1_packed_bits_to_float, float_to_e2m1_bits};
fn decode(code: u32) -> f32 {
e2m1::from_bits((code & 0xF) as u8).to_f32()
}
fn encode(value: f32) -> u32 {
e2m1::from_f32(value).to_bits() as u32
}
#[cube(launch_unchecked)]
fn kernel_decode<N: Size>(codes: &[Vector<u32, N>], out: &mut [Vector<f32, N>]) {
if ABSOLUTE_POS < codes.len() {
out[ABSOLUTE_POS] = e2m1_bits_to_float::<f32, N>(codes[ABSOLUTE_POS]);
}
}
#[cube(launch_unchecked)]
fn kernel_encode<N: Size>(values: &[Vector<f32, N>], out: &mut [Vector<u32, N>]) {
if ABSOLUTE_POS < values.len() {
out[ABSOLUTE_POS] = float_to_e2m1_bits::<f32, N>(values[ABSOLUTE_POS]);
}
}
#[cube(launch_unchecked)]
fn kernel_decode_packed<N: Size>(words: &[u32], out: &mut [Vector<f32, N>]) {
if ABSOLUTE_POS < words.len() {
out[ABSOLUTE_POS] = e2m1_packed_bits_to_float::<f32, N>(words[ABSOLUTE_POS]);
}
}
pub fn test_e2m1_codec_matches_host<R: Runtime>(client: Client) {
let codes: Vec<u32> = (0..16u32).chain((0..16u32).map(|c| c | 0xFFF0)).collect();
let decoded = launch_decode(&client, &codes);
let mut bad = vec![];
for (i, &code) in codes.iter().enumerate() {
let expected = decode(code);
if decoded[i].to_bits() != expected.to_bits() {
bad.push(format!(
" decode {code:#x}: device {}, host {expected}",
decoded[i]
));
}
}
let values = encode_inputs();
let encoded = launch_encode(&client, &values);
for (i, &value) in values.iter().enumerate() {
if value.is_nan() {
assert!(
encoded[i] <= 0xF,
"encode NaN left the nibble: {:#x}",
encoded[i]
);
continue;
}
let expected = encode(value);
if encoded[i] != expected {
bad.push(format!(
" encode {value:e}: device {:#x}, host {expected:#x}",
encoded[i]
));
}
}
let words: Vec<u32> = (0..=0xFFu32).collect();
let packed = launch_decode_packed(&client, &words);
for (i, &word) in words.iter().enumerate() {
for lane in 0..2 {
let expected = decode(word >> (4 * lane));
let actual = packed[2 * i + lane as usize];
if actual.to_bits() != expected.to_bits() {
bad.push(format!(
" packed {word:#04x} lane {lane}: device {actual}, host {expected}"
));
}
}
}
if !bad.is_empty() {
panic!(
"{} disagree with the host codec\n{}",
bad.len(),
bad.join("\n")
);
}
}
fn encode_inputs() -> Vec<f32> {
let mut values = vec![];
for code in 0..16u32 {
values.push(decode(code));
}
for midpoint in [0.25f32, 0.75, 1.25, 1.75, 2.5, 3.5, 5.0] {
for value in [midpoint, midpoint * 0.999, midpoint * 1.001] {
values.push(value);
values.push(-value);
}
}
for value in [6.1f32, 100.0, f32::MAX, f32::INFINITY] {
values.push(value);
values.push(-value);
}
values.push(f32::NAN);
values
}
fn launch_decode(client: &Client, codes: &[u32]) -> Vec<f32> {
let handle_in = client.create_from_slice(u32::as_bytes(codes));
let handle_out = client.empty(size_of_val(codes));
let cube_dim = 256u32;
unsafe {
kernel_decode::launch_unchecked(
client,
CubeCount::Static((codes.len() as u32).div_ceil(cube_dim), 1, 1),
CubeDim::new_1d(cube_dim),
1,
BufferArg::from_raw_parts(handle_in, codes.len()),
BufferArg::from_raw_parts(handle_out.clone(), codes.len()),
)
};
f32::from_bytes(&client.read_one_unchecked(handle_out)).to_vec()
}
fn launch_encode(client: &Client, values: &[f32]) -> Vec<u32> {
let handle_in = client.create_from_slice(f32::as_bytes(values));
let handle_out = client.empty(size_of_val(values));
let cube_dim = 256u32;
unsafe {
kernel_encode::launch_unchecked(
client,
CubeCount::Static((values.len() as u32).div_ceil(cube_dim), 1, 1),
CubeDim::new_1d(cube_dim),
1,
BufferArg::from_raw_parts(handle_in, values.len()),
BufferArg::from_raw_parts(handle_out.clone(), values.len()),
)
};
u32::from_bytes(&client.read_one_unchecked(handle_out)).to_vec()
}
fn launch_decode_packed(client: &Client, words: &[u32]) -> Vec<f32> {
let lanes = 2;
let handle_in = client.create_from_slice(u32::as_bytes(words));
let handle_out = client.empty(words.len() * lanes * size_of::<f32>());
let cube_dim = 256u32;
unsafe {
kernel_decode_packed::launch_unchecked(
client,
CubeCount::Static((words.len() as u32).div_ceil(cube_dim), 1, 1),
CubeDim::new_1d(cube_dim),
lanes,
BufferArg::from_raw_parts(handle_in, words.len()),
BufferArg::from_raw_parts(handle_out.clone(), words.len() * lanes),
)
};
f32::from_bytes(&client.read_one_unchecked(handle_out)).to_vec()
}
#[allow(missing_docs)]
#[macro_export]
macro_rules! testgen_fp4 {
() => {
mod fp4 {
use super::*;
#[$crate::tests::test_log::test]
fn e2m1_codec_matches_host() {
let client = TestRuntime::client(&Default::default());
cubecl_std::tests::fp4::test_e2m1_codec_matches_host::<TestRuntime>(client);
}
}
};
}