Skip to main content

virtio_accel_tosa/
integer.rs

1//! Exact integer primitives shared by TOSA conformance and provider legalization.
2//!
3//! These functions implement the load-bearing arithmetic from TOSA 1.0.1 rather than relying on
4//! Rust's build-profile-dependent overflow behavior. A provider may use wider native operations,
5//! but its visible result must match these helpers for every predictable input.
6
7/// A TOSA integer arithmetic precondition was not satisfied.
8#[derive(Clone, Copy, Debug, PartialEq, Eq)]
9pub enum IntegerError {
10    /// Operands that form one dot product have different lengths.
11    LengthMismatch,
12    /// TOSA requires a non-negative fixed-point multiplier.
13    NegativeMultiplier,
14    /// TOSA scaling shifts are restricted to the inclusive range `2..=62`.
15    ShiftOutOfRange,
16    /// The input violates an operator `REQUIRE` range.
17    InputOutOfRange,
18    /// Exact intermediate arithmetic does not fit its TOSA accumulator type.
19    AccumulatorOverflow,
20}
21
22/// Compute an exact signed INT8 dot product with INT32 accumulation.
23///
24/// `left` and `right` contain the tensor's two's-complement storage bytes. Zero points are
25/// subtracted before multiplication, and `bias` initializes the accumulator. Wrapping is never
26/// used: a result outside INT32 is reported because TOSA classifies such an input as unpredictable
27/// rather than defining modular arithmetic.
28pub fn dot_i8_i32(
29    left: &[u8],
30    right: &[u8],
31    left_zero_point: i8,
32    right_zero_point: i8,
33    bias: i32,
34) -> Result<i32, IntegerError> {
35    if left.len() != right.len() {
36        return Err(IntegerError::LengthMismatch);
37    }
38    let left_zero_point = i64::from(left_zero_point);
39    let right_zero_point = i64::from(right_zero_point);
40    let mut accumulator = i64::from(bias);
41    for (&left, &right) in left.iter().zip(right) {
42        let left = i64::from(left as i8) - left_zero_point;
43        let right = i64::from(right as i8) - right_zero_point;
44        accumulator = accumulator
45            .checked_add(left * right)
46            .ok_or(IntegerError::AccumulatorOverflow)?;
47    }
48    i32::try_from(accumulator).map_err(|_| IntegerError::AccumulatorOverflow)
49}
50
51/// Apply TOSA's 32-bit fixed-point scaling helper exactly.
52///
53/// This is `apply_scale_32` from TOSA 1.0.1. The arithmetic right shift is explicit and every
54/// pseudocode `REQUIRE` is checked before calculating the result.
55pub fn apply_scale_32(
56    value: i32,
57    multiplier: i32,
58    shift: i8,
59    double_round: bool,
60) -> Result<i32, IntegerError> {
61    if multiplier < 0 {
62        return Err(IntegerError::NegativeMultiplier);
63    }
64    let shift = checked_shift(shift)?;
65    let value64 = i64::from(value);
66    let bound = 1_i64 << (shift - 1);
67    if !(-bound..bound).contains(&value64) {
68        return Err(IntegerError::InputOutOfRange);
69    }
70
71    let mut round = 1_i64 << (shift - 1);
72    if double_round && shift > 31 {
73        if value >= 0 {
74            round += 1_i64 << 30;
75        } else {
76            round -= 1_i64 << 30;
77        }
78    }
79    let result = value64
80        .checked_mul(i64::from(multiplier))
81        .and_then(|value| value.checked_add(round))
82        .ok_or(IntegerError::AccumulatorOverflow)?
83        >> shift;
84    i32::try_from(result).map_err(|_| IntegerError::AccumulatorOverflow)
85}
86
87/// Apply TOSA's 16-bit fixed-point scaling helper exactly.
88///
89/// `value` is represented in an `i64`, but must fit TOSA's signed 48-bit accumulator domain.
90pub fn apply_scale_16(value: i64, multiplier: i16, shift: i8) -> Result<i32, IntegerError> {
91    if multiplier < 0 {
92        return Err(IntegerError::NegativeMultiplier);
93    }
94    if !(-(1_i64 << 47)..(1_i64 << 47)).contains(&value) {
95        return Err(IntegerError::InputOutOfRange);
96    }
97    let shift = checked_shift(shift)?;
98    let round = 1_i64 << (shift - 1);
99    let result = value
100        .checked_mul(i64::from(multiplier))
101        .and_then(|value| value.checked_add(round))
102        .ok_or(IntegerError::AccumulatorOverflow)?
103        >> shift;
104    i32::try_from(result).map_err(|_| IntegerError::AccumulatorOverflow)
105}
106
107/// Rescale one INT32 accumulator into a signed INT8 tensor value.
108///
109/// The input zero point for an INT32 TOSA tensor is necessarily zero. `output_zero_point` is added
110/// after scaling and the final value is saturated to the signed INT8 storage range.
111pub fn rescale_i32_to_i8(
112    value: i32,
113    multiplier: i32,
114    shift: i8,
115    output_zero_point: i8,
116    double_round: bool,
117) -> Result<i8, IntegerError> {
118    let scaled = apply_scale_32(value, multiplier, shift, double_round)?;
119    let shifted = i64::from(scaled) + i64::from(output_zero_point);
120    Ok(shifted.clamp(i64::from(i8::MIN), i64::from(i8::MAX)) as i8)
121}
122
123const fn checked_shift(shift: i8) -> Result<u32, IntegerError> {
124    if shift < 2 || shift > 62 {
125        Err(IntegerError::ShiftOutOfRange)
126    } else {
127        Ok(shift as u32)
128    }
129}
130
131#[cfg(test)]
132mod tests {
133    use super::*;
134
135    #[test]
136    fn dot_product_interprets_storage_as_signed_and_applies_zero_points() {
137        let left = [0x80, 0xff, 0x00, 0x7f];
138        let right = [0x7f, 0x01, 0xff, 0x80];
139        // Sum of each `(left - -2) * (right - 1)` term plus the bias.
140        let expected = -32_504;
141        assert_eq!(dot_i8_i32(&left, &right, -2, 1, 17), Ok(expected));
142    }
143
144    #[test]
145    fn dot_product_rejects_shape_and_accumulator_violations() {
146        assert_eq!(
147            dot_i8_i32(&[1], &[], 0, 0, 0),
148            Err(IntegerError::LengthMismatch)
149        );
150        let positive = alloc::vec![0x7f; 140_000];
151        assert_eq!(
152            dot_i8_i32(&positive, &positive, -128, -128, i32::MAX),
153            Err(IntegerError::AccumulatorOverflow)
154        );
155    }
156
157    #[test]
158    fn scale_32_matches_tosa_unity_and_signed_rounding() {
159        let unity = 1_i32 << 30;
160        for value in [-127, -1, 0, 1, 127] {
161            assert_eq!(apply_scale_32(value, unity, 30, false), Ok(value));
162        }
163        assert_eq!(apply_scale_32(3, 2, 3, false), Ok(1));
164        assert_eq!(apply_scale_32(-3, 2, 3, false), Ok(-1));
165    }
166
167    #[test]
168    fn scale_32_checks_every_pseudocode_precondition() {
169        assert_eq!(
170            apply_scale_32(0, -1, 30, false),
171            Err(IntegerError::NegativeMultiplier)
172        );
173        assert_eq!(
174            apply_scale_32(0, 1, 1, false),
175            Err(IntegerError::ShiftOutOfRange)
176        );
177        assert_eq!(
178            apply_scale_32(2, 1, 2, false),
179            Err(IntegerError::InputOutOfRange)
180        );
181    }
182
183    #[test]
184    fn scale_16_matches_tosa_unity_and_checks_int48() {
185        let unity = 1_i16 << 14;
186        for value in [-32_768, -1, 0, 1, 32_767] {
187            assert_eq!(apply_scale_16(value, unity, 14), Ok(value as i32));
188        }
189        assert_eq!(
190            apply_scale_16(1_i64 << 47, unity, 14),
191            Err(IntegerError::InputOutOfRange)
192        );
193    }
194
195    #[test]
196    fn rescale_saturates_only_after_output_zero_point() {
197        let unity = 1_i32 << 30;
198        assert_eq!(rescale_i32_to_i8(100, unity, 30, 20, false), Ok(120));
199        assert_eq!(rescale_i32_to_i8(120, unity, 30, 20, false), Ok(127));
200        assert_eq!(rescale_i32_to_i8(-120, unity, 30, -20, false), Ok(-128));
201    }
202}