virtio_accel_tosa/
numeric.rs1use crate::DType;
7
8pub const fn low_precision_storage_bytes(dtype: DType, elements: usize) -> Option<usize> {
13 match dtype {
14 DType::INT4 => Some(elements / 2 + elements % 2),
15 DType::INT8 | DType::FP8E4M3 | DType::FP8E5M2 => Some(elements),
16 _ => None,
17 }
18}
19
20pub fn unpack_int4(bytes: &[u8], index: usize) -> Option<i8> {
25 let byte = *bytes.get(index / 2)?;
26 let nibble = if index % 2 == 0 {
27 byte & 0x0f
28 } else {
29 byte >> 4
30 };
31 Some((nibble as i8) << 4 >> 4)
32}
33
34pub const fn pack_int4(low: i8, high: i8) -> Option<u8> {
38 if low < -7 || low > 7 || high < -7 || high > 7 {
39 return None;
40 }
41 Some(((high as u8 & 0x0f) << 4) | (low as u8 & 0x0f))
42}
43
44pub fn fp8e4m3_to_f32(bits: u8) -> f32 {
49 let sign = u32::from(bits & 0x80) << 24;
50 let exponent = u32::from((bits >> 3) & 0x0f);
51 let fraction = u32::from(bits & 0x07);
52 if exponent == 0 {
53 return signed_fp8_subnormal(sign, fraction, 9);
54 }
55 if exponent == 0x0f && fraction == 0x07 {
56 return f32::from_bits(sign | 0x7fc0_0000);
57 }
58 f32::from_bits(sign | ((exponent + 120) << 23) | (fraction << 20))
59}
60
61pub fn fp8e5m2_to_f32(bits: u8) -> f32 {
63 let sign = u32::from(bits & 0x80) << 24;
64 let exponent = u32::from((bits >> 2) & 0x1f);
65 let fraction = u32::from(bits & 0x03);
66 if exponent == 0 {
67 return signed_fp8_subnormal(sign, fraction, 16);
68 }
69 if exponent == 0x1f {
70 return f32::from_bits(sign | 0x7f80_0000 | (fraction << 21));
71 }
72 f32::from_bits(sign | ((exponent + 112) << 23) | (fraction << 21))
73}
74
75pub const fn fp8e4m3_to_bf16_bits(bits: u8) -> u16 {
80 let sign = ((bits & 0x80) as u16) << 8;
81 let exponent = ((bits >> 3) & 0x0f) as u16;
82 let fraction = (bits & 0x07) as u16;
83 if exponent == 0 {
84 let subnormal = [
85 0x0000, 0x3b00, 0x3b80, 0x3bc0, 0x3c00, 0x3c20, 0x3c40, 0x3c60,
86 ];
87 return sign | subnormal[fraction as usize];
88 }
89 if exponent == 0x0f && fraction == 0x07 {
90 return sign | 0x7fc0;
91 }
92 sign | ((exponent + 120) << 7) | (fraction << 4)
93}
94
95pub const fn fp8e5m2_to_bf16_bits(bits: u8) -> u16 {
100 let sign = ((bits & 0x80) as u16) << 8;
101 let exponent = ((bits >> 2) & 0x1f) as u16;
102 let fraction = (bits & 0x03) as u16;
103 if exponent == 0 {
104 let subnormal = [0x0000, 0x3780, 0x3800, 0x3840];
105 return sign | subnormal[fraction as usize];
106 }
107 if exponent == 0x1f {
108 return sign | if fraction == 0 { 0x7f80 } else { 0x7fc0 };
109 }
110 sign | ((exponent + 112) << 7) | (fraction << 5)
111}
112
113fn signed_fp8_subnormal(sign: u32, fraction: u32, scale: i32) -> f32 {
114 if fraction == 0 {
115 return f32::from_bits(sign);
116 }
117 let power_of_two = f32::from_bits(((127 - scale) as u32) << 23);
118 let value = (fraction as f32) * power_of_two;
119 f32::from_bits(value.to_bits() | sign)
120}
121
122#[cfg(test)]
123mod tests {
124 use super::*;
125
126 #[test]
127 fn low_precision_sizes_are_checked_and_int4_is_nibble_packed() {
128 assert_eq!(low_precision_storage_bytes(DType::INT4, 0), Some(0));
129 assert_eq!(low_precision_storage_bytes(DType::INT4, 7), Some(4));
130 assert_eq!(low_precision_storage_bytes(DType::INT4, 8), Some(4));
131 assert_eq!(
132 low_precision_storage_bytes(DType::INT4, usize::MAX),
133 Some(usize::MAX / 2 + 1)
134 );
135 assert_eq!(low_precision_storage_bytes(DType::INT8, 8), Some(8));
136 assert_eq!(low_precision_storage_bytes(DType::FP8E4M3, 8), Some(8));
137 assert_eq!(low_precision_storage_bytes(DType::FP8E5M2, 8), Some(8));
138 assert_eq!(low_precision_storage_bytes(DType::FP16, 8), None);
139
140 assert_eq!(pack_int4(-7, 7), Some(0x79));
141 assert_eq!(pack_int4(-8, 0), None);
142 assert_eq!(unpack_int4(&[0x79], 0), Some(-7));
143 assert_eq!(unpack_int4(&[0x79], 1), Some(7));
144 assert_eq!(unpack_int4(&[0x08], 0), Some(-8));
145 assert_eq!(unpack_int4(&[], 0), None);
146 }
147
148 #[test]
149 fn fp8e4m3_decodes_signed_zero_subnormal_finite_max_and_nan() {
150 assert_eq!(fp8e4m3_to_f32(0x00).to_bits(), 0.0_f32.to_bits());
151 assert_eq!(fp8e4m3_to_f32(0x80).to_bits(), (-0.0_f32).to_bits());
152 assert_eq!(fp8e4m3_to_f32(0x01), 2.0_f32.powi(-9));
153 assert_eq!(fp8e4m3_to_f32(0x38), 1.0);
154 assert_eq!(fp8e4m3_to_f32(0x7e), 448.0);
155 assert!(fp8e4m3_to_f32(0x7f).is_nan());
156 assert!(fp8e4m3_to_f32(0xff).is_nan());
157 }
158
159 #[test]
160 fn fp8e5m2_decodes_signed_zero_subnormal_finite_max_infinity_and_nan() {
161 assert_eq!(fp8e5m2_to_f32(0x00).to_bits(), 0.0_f32.to_bits());
162 assert_eq!(fp8e5m2_to_f32(0x80).to_bits(), (-0.0_f32).to_bits());
163 assert_eq!(fp8e5m2_to_f32(0x01), 2.0_f32.powi(-16));
164 assert_eq!(fp8e5m2_to_f32(0x3c), 1.0);
165 assert_eq!(fp8e5m2_to_f32(0x7b), 57_344.0);
166 assert_eq!(fp8e5m2_to_f32(0x7c), f32::INFINITY);
167 assert_eq!(fp8e5m2_to_f32(0xfc), f32::NEG_INFINITY);
168 assert!(fp8e5m2_to_f32(0x7d).is_nan());
169 }
170
171 #[test]
172 fn fp8_to_bf16_is_exact_for_every_non_nan_encoding() {
173 for bits in u8::MIN..=u8::MAX {
174 for (value, bf16) in [
175 (fp8e4m3_to_f32(bits), fp8e4m3_to_bf16_bits(bits)),
176 (fp8e5m2_to_f32(bits), fp8e5m2_to_bf16_bits(bits)),
177 ] {
178 if value.is_nan() {
179 assert_eq!(bf16 & 0x7f80, 0x7f80);
180 assert_ne!(bf16 & 0x007f, 0);
181 } else {
182 assert_eq!(bf16, (value.to_bits() >> 16) as u16);
183 }
184 }
185 }
186 }
187}