Skip to main content

embedded_nn/
basic_math.rs

1//! Basic elementwise mathematical operations on quantized tensors.
2
3use crate::support::{clamp, requantize};
4use crate::types::{Activation, Result};
5
6/// Elementwise addition parameters for quantized int8 tensors.
7#[derive(Debug, Clone, Copy, PartialEq, Eq)]
8pub struct ElementwiseAddParams {
9    /// Zero point offset for input 1.
10    pub input1_offset: i32,
11    /// Multiplier for input 1.
12    pub input1_mult: i32,
13    /// Shift for input 1.
14    pub input1_shift: i32,
15    /// Zero point offset for input 2.
16    pub input2_offset: i32,
17    /// Multiplier for input 2.
18    pub input2_mult: i32,
19    /// Shift for input 2.
20    pub input2_shift: i32,
21    /// Common left shift for inputs.
22    pub left_shift: i32,
23    /// Output zero point offset.
24    pub output_offset: i32,
25    /// Output multiplier.
26    pub output_mult: i32,
27    /// Output shift.
28    pub output_shift: i32,
29    /// Output activation range.
30    pub activation: Activation,
31}
32
33/// Elementwise multiplication parameters for quantized int8 tensors.
34#[derive(Debug, Clone, Copy, PartialEq, Eq)]
35pub struct ElementwiseMulParams {
36    /// Zero point offset for input 1.
37    pub input1_offset: i32,
38    /// Zero point offset for input 2.
39    pub input2_offset: i32,
40    /// Output zero point offset.
41    pub output_offset: i32,
42    /// Output multiplier.
43    pub output_mult: i32,
44    /// Output shift.
45    pub output_shift: i32,
46    /// Output activation range.
47    pub activation: Activation,
48}
49
50/// Performs elementwise addition of two int8 tensors.
51pub fn elementwise_add_s8(
52    input1: &[i8],
53    input2: &[i8],
54    output: &mut [i8],
55    params: &ElementwiseAddParams,
56) -> Result<()> {
57    let size = input1.len().min(input2.len()).min(output.len());
58    for i in 0..size {
59        let val1 = (input1[i] as i32 + params.input1_offset) << params.left_shift;
60        let val2 = (input2[i] as i32 + params.input2_offset) << params.left_shift;
61
62        let req1 = requantize(val1, params.input1_mult, params.input1_shift);
63        let req2 = requantize(val2, params.input2_mult, params.input2_shift);
64
65        let sum = req1 + req2;
66        let req_sum = requantize(sum, params.output_mult, params.output_shift);
67        let final_val = req_sum + params.output_offset;
68
69        output[i] = clamp(final_val, params.activation.min, params.activation.max) as i8;
70    }
71    Ok(())
72}
73
74/// Performs elementwise multiplication of two int8 tensors.
75pub fn elementwise_mul_s8(
76    input1: &[i8],
77    input2: &[i8],
78    output: &mut [i8],
79    params: &ElementwiseMulParams,
80) -> Result<()> {
81    let size = input1.len().min(input2.len()).min(output.len());
82    for i in 0..size {
83        let val1 = input1[i] as i32 + params.input1_offset;
84        let val2 = input2[i] as i32 + params.input2_offset;
85        let prod = val1 * val2;
86
87        let req_prod = requantize(prod, params.output_mult, params.output_shift);
88        let final_val = req_prod + params.output_offset;
89
90        output[i] = clamp(final_val, params.activation.min, params.activation.max) as i8;
91    }
92    Ok(())
93}
94
95/// Performs elementwise subtraction of two int8 tensors.
96pub fn elementwise_sub_s8(
97    input1: &[i8],
98    input2: &[i8],
99    output: &mut [i8],
100    params: &ElementwiseAddParams,
101) -> Result<()> {
102    let size = input1.len().min(input2.len()).min(output.len());
103    for i in 0..size {
104        let val1 = (input1[i] as i32 + params.input1_offset) << params.left_shift;
105        let val2 = (input2[i] as i32 + params.input2_offset) << params.left_shift;
106
107        let req1 = requantize(val1, params.input1_mult, params.input1_shift);
108        let req2 = requantize(val2, params.input2_mult, params.input2_shift);
109
110        let diff = req1 - req2;
111        let req_diff = requantize(diff, params.output_mult, params.output_shift);
112        let final_val = req_diff + params.output_offset;
113
114        output[i] = clamp(final_val, params.activation.min, params.activation.max) as i8;
115    }
116    Ok(())
117}
118
119/// Performs elementwise addition of two int16 tensors.
120pub fn elementwise_add_s16(
121    input1: &[i16],
122    input2: &[i16],
123    output: &mut [i16],
124    mult1: i32,
125    shift1: i32,
126    mult2: i32,
127    shift2: i32,
128    output_mult: i32,
129    output_shift: i32,
130    act: Activation,
131) -> Result<()> {
132    let size = input1.len().min(input2.len()).min(output.len());
133    for i in 0..size {
134        let req1 = requantize(input1[i] as i32, mult1, shift1);
135        let req2 = requantize(input2[i] as i32, mult2, shift2);
136
137        let sum = req1 + req2;
138        let req_sum = requantize(sum, output_mult, output_shift);
139
140        output[i] = clamp(req_sum, act.min, act.max) as i16;
141    }
142    Ok(())
143}
144
145/// Performs elementwise multiplication of two int16 tensors.
146pub fn elementwise_mul_s16(
147    input1: &[i16],
148    input2: &[i16],
149    output: &mut [i16],
150    output_mult: i32,
151    output_shift: i32,
152    act: Activation,
153) -> Result<()> {
154    let size = input1.len().min(input2.len()).min(output.len());
155    for i in 0..size {
156        let prod = (input1[i] as i64) * (input2[i] as i64);
157        let req_prod = requantize((prod >> 15) as i32, output_mult, output_shift);
158        output[i] = clamp(req_prod, act.min, act.max) as i16;
159    }
160    Ok(())
161}
162
163#[cfg(test)]
164mod tests {
165    use super::*;
166
167    #[test]
168    fn test_elementwise_add_s8() {
169        let input1 = [10i8, 20i8, 30i8];
170        let input2 = [5i8, 15i8, 25i8];
171        let mut output = [0i8; 3];
172
173        let params = ElementwiseAddParams {
174            input1_offset: 0,
175            input1_mult: 1073741824, // 0.5
176            input1_shift: 0,
177            input2_offset: 0,
178            input2_mult: 1073741824, // 0.5
179            input2_shift: 0,
180            left_shift: 0,
181            output_offset: 0,
182            output_mult: 1073741824, // 0.5
183            output_shift: 0,
184            activation: Activation::int8_unconstrained(),
185        };
186
187        elementwise_add_s8(&input1, &input2, &mut output, &params).unwrap();
188        // (10*0.5 + 5*0.5) * 0.5 = 3.75 -> rounded to 4
189        assert!((output[0] - 4).abs() <= 1);
190    }
191}