Skip to main content

virtio_accel_tosa/
numeric.rs

1//! Stable-Rust helpers for TOSA low-precision tensor encodings.
2//!
3//! These functions describe the device-neutral wire representation. They do not imply that a
4//! particular accelerator can execute a graph containing the corresponding [`DType`].
5
6use crate::DType;
7
8/// Return the packed byte count for an INT4, INT8, FP8E4M3, or FP8E5M2 tensor.
9///
10/// INT4 elements are packed low nibble first. `None` means that `dtype` is not one of these
11/// low-precision types. The packed byte count itself cannot overflow `usize`.
12pub 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
20/// Decode one signed two's-complement INT4 value packed low nibble first.
21///
22/// This returns all mechanically representable values, including `-8`. TOSA 1.0 tensor constants
23/// use the narrower `-7..=7` range; semantic validation rejects `-8` where the specification does.
24pub 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
34/// Pack two TOSA INT4 values, with `low` in the low nibble and `high` in the high nibble.
35///
36/// TOSA 1.0 defines INT4 values over `-7..=7`, so `-8` and out-of-range inputs are rejected.
37pub 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
44/// Convert one TOSA/OCP FP8 E4M3 bit pattern to an exactly representable `f32` value.
45///
46/// E4M3 has no infinities. Exponent `0b1111` remains finite except when the fraction is `0b111`,
47/// which represents NaN.
48pub 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
61/// Convert one TOSA/OCP FP8 E5M2 bit pattern to an exactly representable `f32` value.
62pub 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
75/// Convert one TOSA/OCP FP8 E4M3 bit pattern to BF16 storage bits.
76///
77/// Every finite E4M3 value is exactly representable in BF16. Signed zero is preserved and NaN is
78/// canonicalized to a quiet BF16 NaN while preserving its sign.
79pub 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
95/// Convert one TOSA/OCP FP8 E5M2 bit pattern to BF16 storage bits.
96///
97/// Every finite E5M2 value and infinity is exactly representable in BF16. Signed zero is
98/// preserved and NaN is canonicalized to a quiet BF16 NaN while preserving its sign.
99pub 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}