1#[derive(Clone, Copy, Debug, PartialEq, Eq)]
9pub enum IntegerError {
10 LengthMismatch,
12 NegativeMultiplier,
14 ShiftOutOfRange,
16 InputOutOfRange,
18 AccumulatorOverflow,
20}
21
22pub 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
51pub 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
87pub 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
107pub 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 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}