1use anyhow::{Result, ensure};
16
17pub const E2M1: [f32; 16] = [
19 0.0, 0.5, 1.0, 1.5, 2.0, 3.0, 4.0, 6.0, -0.0, -0.5, -1.0, -1.5, -2.0, -3.0, -4.0, -6.0,
20];
21
22pub const GROUP_SIZE: usize = 32;
24
25const FP8_E8M0_LUT: [f32; 256] = {
27 let mut table = [0.0f32; 256];
28 let mut i: u32 = 0;
29 while i < 256 {
30 let exp = i as u8;
31 table[i as usize] = if exp == 0 || exp == 255 {
32 0.0f32
33 } else {
34 f32::from_bits((exp as u32) << 23)
35 };
36 i += 1;
37 }
38 table
39};
40
41#[inline(always)]
43pub fn fp8_e8m0_to_f32(bits: u8) -> f32 {
44 FP8_E8M0_LUT[bits as usize]
45}
46
47pub fn dequant_nvfp4_e8m0_to_f32(
52 packed: &[u8],
53 scales: &[u8],
54 n: usize,
55 k: usize,
56) -> Result<Vec<f32>> {
57 let total = n.checked_mul(k).expect("mxfp4 n*k");
58 ensure!(
59 total.is_multiple_of(2),
60 "MXFP4 E8M0: n*k={total} is odd (need even nibble count)"
61 );
62 let packed_bytes = total / 2;
63 ensure!(
64 packed.len() == packed_bytes,
65 "MXFP4 E8M0: packed {} B, expected {packed_bytes} for [{n},{k}]",
66 packed.len()
67 );
68 let num_groups = scales.len();
69 ensure!(
70 num_groups > 0 && total.is_multiple_of(num_groups),
71 "MXFP4 E8M0: weight elems {total} not divisible by E8M0 scale groups {num_groups}"
72 );
73 let block = total / num_groups;
74 let mut out = vec![0.0f32; total];
75 for (group, &sb) in scales.iter().enumerate() {
76 let block_scale = fp8_e8m0_to_f32(sb);
77 for elem in 0..block {
78 let flat_idx = group * block + elem;
79 let byte_idx = flat_idx / 2;
80 let nibble = if flat_idx.is_multiple_of(2) {
81 packed[byte_idx] & 0x0F
82 } else {
83 (packed[byte_idx] >> 4) & 0x0F
84 };
85 out[flat_idx] = E2M1[nibble as usize] * block_scale;
86 }
87 }
88 Ok(out)
89}
90
91pub fn dequant_nvfp4_e8m0_to_bf16(
93 packed: &[u8],
94 scales: &[u8],
95 n: usize,
96 k: usize,
97) -> Result<Vec<u16>> {
98 let f = dequant_nvfp4_e8m0_to_f32(packed, scales, n, k)?;
99 Ok(f.into_iter().map(crate::numeric::f32_to_bf16).collect())
100}
101
102#[cfg(test)]
103mod tests {
104 use super::*;
105
106 #[test]
107 fn e8m0_pow2_and_sentinels() {
108 assert_eq!(fp8_e8m0_to_f32(0), 0.0);
109 assert_eq!(fp8_e8m0_to_f32(255), 0.0);
110 assert_eq!(fp8_e8m0_to_f32(127), 1.0);
111 assert_eq!(fp8_e8m0_to_f32(128), 2.0);
112 assert_eq!(fp8_e8m0_to_f32(126), 0.5);
113 }
114
115 #[test]
116 fn matches_dsv4_nibble_order_and_lut() {
117 let packed = [0x12u8]; let scales = [127u8]; let got = dequant_nvfp4_e8m0_to_f32(&packed, &scales, 1, 2).unwrap();
121 assert_eq!(got, vec![1.0, 0.5]);
122 }
123}