1use crate::support::{clamp, requantize};
4use crate::types::Activation;
5
6pub const SIGMOID_TABLE_UINT16: [u16; 256] = [
8 32768, 33451, 34133, 34813, 35493, 36169, 36843, 37513, 38180, 38841, 39498, 40149, 40794,
9 41432, 42064, 42688, 43304, 43912, 44511, 45102, 45683, 46255, 46817, 47369, 47911, 48443,
10 48964, 49475, 49975, 50464, 50942, 51409, 51865, 52311, 52745, 53169, 53581, 53983, 54374,
11 54755, 55125, 55485, 55834, 56174, 56503, 56823, 57133, 57433, 57724, 58007, 58280, 58544,
12 58800, 59048, 59288, 59519, 59743, 59959, 60168, 60370, 60565, 60753, 60935, 61110, 61279,
13 61441, 61599, 61750, 61896, 62036, 62172, 62302, 62428, 62549, 62666, 62778, 62886, 62990,
14 63090, 63186, 63279, 63368, 63454, 63536, 63615, 63691, 63765, 63835, 63903, 63968, 64030,
15 64090, 64148, 64204, 64257, 64308, 64357, 64405, 64450, 64494, 64536, 64576, 64614, 64652,
16 64687, 64721, 64754, 64786, 64816, 64845, 64873, 64900, 64926, 64950, 64974, 64997, 65019,
17 65039, 65060, 65079, 65097, 65115, 65132, 65149, 65164, 65179, 65194, 65208, 65221, 65234,
18 65246, 65258, 65269, 65280, 65291, 65301, 65310, 65319, 65328, 65337, 65345, 65352, 65360,
19 65367, 65374, 65381, 65387, 65393, 65399, 65404, 65410, 65415, 65420, 65425, 65429, 65433,
20 65438, 65442, 65445, 65449, 65453, 65456, 65459, 65462, 65465, 65468, 65471, 65474, 65476,
21 65479, 65481, 65483, 65485, 65488, 65489, 65491, 65493, 65495, 65497, 65498, 65500, 65501,
22 65503, 65504, 65505, 65507, 65508, 65509, 65510, 65511, 65512, 65513, 65514, 65515, 65516,
23 65517, 65517, 65518, 65519, 65520, 65520, 65521, 65522, 65522, 65523, 65523, 65524, 65524,
24 65525, 65525, 65526, 65526, 65526, 65527, 65527, 65528, 65528, 65528, 65529, 65529, 65529,
25 65529, 65530, 65530, 65530, 65530, 65531, 65531, 65531, 65531, 65531, 65532, 65532, 65532,
26 65532, 65532, 65532, 65533, 65533, 65533, 65533, 65533, 65533, 65533, 65533, 65534, 65534,
27 65534, 65534, 65534, 65534, 65534, 65534, 65534, 65534, 65535,
28];
29
30pub fn relu_s8(data: &mut [i8]) {
32 for val in data.iter_mut() {
33 if *val < 0 {
34 *val = 0;
35 }
36 }
37}
38
39pub fn relu6_s8(data: &mut [i8]) {
41 for val in data.iter_mut() {
42 let mut ip = *val as i32;
43 if ip < 0 {
44 ip = 0;
45 }
46 if ip > 6 {
47 ip = 6;
48 }
49 *val = ip as i8;
50 }
51}
52
53pub fn activation_s8(data: &mut [i8], act: Activation) {
55 for val in data.iter_mut() {
56 let clamped = clamp(*val as i32, act.min, act.max);
57 *val = clamped as i8;
58 }
59}
60
61pub fn relu_s16(data: &mut [i16]) {
63 for val in data.iter_mut() {
64 if *val < 0 {
65 *val = 0;
66 }
67 }
68}
69
70pub fn activation_s16(data: &mut [i16], act: Activation) {
72 for val in data.iter_mut() {
73 let clamped = clamp(*val as i32, act.min, act.max);
74 *val = clamped as i16;
75 }
76}
77
78pub fn leaky_relu_s8(
80 input: &[i8],
81 output: &mut [i8],
82 alpha_mult: i32,
83 alpha_shift: i32,
84 input_offset: i32,
85 output_offset: i32,
86) {
87 let size = input.len().min(output.len());
88 for i in 0..size {
89 let val = input[i] as i32 + input_offset;
90 let res = if val < 0 {
91 requantize(val, alpha_mult, alpha_shift) + output_offset
92 } else {
93 val + output_offset
94 };
95 output[i] = clamp(res, i8::MIN as i32, i8::MAX as i32) as i8;
96 }
97}
98
99pub fn sigmoid_s8(input: &[i8], output: &mut [i8]) {
101 let size = input.len().min(output.len());
102 for i in 0..size {
103 let val = input[i] as i32;
104 let abs_val = val.abs();
105 let idx = clamp(abs_val, 0, 255) as usize;
106 let lut_val = SIGMOID_TABLE_UINT16[idx] as u32;
107
108 let q0_16 = if val >= 0 { lut_val } else { 65535 - lut_val };
109
110 let s8_val = ((q0_16 as i32) >> 8) - 128;
111 output[i] = clamp(s8_val, -128, 127) as i8;
112 }
113}
114
115pub fn tanh_s8(input: &[i8], output: &mut [i8]) {
117 let size = input.len().min(output.len());
118 for i in 0..size {
119 let val = input[i] as i32;
120 let abs_val = val.abs();
121 let idx = clamp(abs_val * 2, 0, 255) as usize;
122 let lut_val = SIGMOID_TABLE_UINT16[idx] as i32;
123 let res = ((lut_val - 32768) * 2) >> 8;
124 let res_signed = if val >= 0 { res } else { -res };
125 output[i] = clamp(res_signed, -128, 127) as i8;
126 }
127}
128
129pub fn sigmoid_s16(input: &[i16], output: &mut [i16], left_shift: i32) {
131 let size = input.len().min(output.len());
132 let abs_input_shift = 9u32;
133 let max_saturation = (0x7FFF << 10) as u32;
134 let input_multiplier = if left_shift < 0 { 3 } else { 3 << left_shift };
135 let abs_left_shift = if left_shift < 0 {
136 -left_shift as u32
137 } else {
138 0
139 };
140 let rounding = if abs_left_shift > 0 {
141 1 << (abs_left_shift - 1)
142 } else {
143 0
144 };
145
146 for i in 0..size {
147 let input_data = ((input[i] as i32) * input_multiplier + rounding) >> abs_left_shift;
148 let abs_input_data = input_data.unsigned_abs();
149 let uh = (abs_input_data >> abs_input_shift) as usize;
150
151 let result = if uh >= 255 {
152 max_saturation
153 } else {
154 let ua = SIGMOID_TABLE_UINT16[uh] as u32;
155 let ub = SIGMOID_TABLE_UINT16[uh + 1] as u32;
156 let ut = abs_input_data & 0x1ff;
157 (ua << abs_input_shift) + ut * (ub - ua)
158 };
159
160 let final_val = if input_data >= 0 {
161 (result + (1 << 9)) >> 10
162 } else {
163 ((1u32 << 25) - result + (1 << 9) - 1) >> 10
164 };
165
166 output[i] = clamp(final_val as i32, i16::MIN as i32, i16::MAX as i32) as i16;
167 }
168}
169
170pub fn tanh_s16(input: &[i16], output: &mut [i16], left_shift: i32) {
172 let size = input.len().min(output.len());
173 let abs_input_shift = 8u32;
174 let max_saturation = (0xFFFF << 8) as u32;
175 let input_multiplier = if left_shift < 0 { 3 } else { 3 << left_shift };
176 let abs_left_shift = if left_shift < 0 {
177 -left_shift as u32
178 } else {
179 0
180 };
181 let rounding = if abs_left_shift > 0 {
182 1 << (abs_left_shift - 1)
183 } else {
184 0
185 };
186
187 for i in 0..size {
188 let input_data = ((input[i] as i32) * input_multiplier + rounding) >> abs_left_shift;
189 let abs_input_data = input_data.unsigned_abs();
190 let uh = (abs_input_data >> abs_input_shift) as usize;
191
192 let result = if uh >= 255 {
193 max_saturation
194 } else {
195 let ua = SIGMOID_TABLE_UINT16[uh] as u32;
196 let ub = SIGMOID_TABLE_UINT16[uh + 1] as u32;
197 let ut = abs_input_data & 0x0ff;
198 (ua << abs_input_shift) + ut * (ub - ua)
199 };
200
201 let pos_val = (((result as i32) - (1 << 23)) + (1 << 7)) >> 8;
202 let final_val = if input_data >= 0 { pos_val } else { -pos_val };
203
204 output[i] = clamp(final_val, i16::MIN as i32, i16::MAX as i32) as i16;
205 }
206}
207
208#[cfg(test)]
209mod tests {
210 use super::*;
211
212 #[test]
213 fn test_relu_s8() {
214 let mut data = [-5i8, 0, 10, -128, 127];
215 relu_s8(&mut data);
216 assert_eq!(data, [0, 0, 10, 0, 127]);
217 }
218
219 #[test]
220 fn test_relu6_s8() {
221 let mut data = [-5i8, 0, 4, 6, 10];
222 relu6_s8(&mut data);
223 assert_eq!(data, [0, 0, 4, 6, 6]);
224 }
225
226 #[test]
227 fn test_sigmoid_s8() {
228 let input = [0i8, 127, -128];
229 let mut output = [0i8; 3];
230 sigmoid_s8(&input, &mut output);
231 assert!((output[0] as i32).abs() <= 5);
233 assert!(output[1] > 100);
234 assert!(output[2] < -100);
235 }
236}