1use crate::support::{clamp, requantize};
4use crate::types::{Activation, Result};
5
6#[derive(Debug, Clone, Copy, PartialEq, Eq)]
8pub struct ElementwiseAddParams {
9 pub input1_offset: i32,
11 pub input1_mult: i32,
13 pub input1_shift: i32,
15 pub input2_offset: i32,
17 pub input2_mult: i32,
19 pub input2_shift: i32,
21 pub left_shift: i32,
23 pub output_offset: i32,
25 pub output_mult: i32,
27 pub output_shift: i32,
29 pub activation: Activation,
31}
32
33#[derive(Debug, Clone, Copy, PartialEq, Eq)]
35pub struct ElementwiseMulParams {
36 pub input1_offset: i32,
38 pub input2_offset: i32,
40 pub output_offset: i32,
42 pub output_mult: i32,
44 pub output_shift: i32,
46 pub activation: Activation,
48}
49
50pub 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
74pub 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
95pub 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
119pub 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
145pub 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, input1_shift: 0,
177 input2_offset: 0,
178 input2_mult: 1073741824, input2_shift: 0,
180 left_shift: 0,
181 output_offset: 0,
182 output_mult: 1073741824, output_shift: 0,
184 activation: Activation::int8_unconstrained(),
185 };
186
187 elementwise_add_s8(&input1, &input2, &mut output, ¶ms).unwrap();
188 assert!((output[0] - 4).abs() <= 1);
190 }
191}