Skip to main content

embedded_nn/
activations.rs

1//! Activation functions for quantized tensors.
2
3use crate::support::{clamp, requantize};
4use crate::types::Activation;
5
6/// Sigmoid and Tanh 256-element lookup table (Q0.16 format).
7pub 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
30/// In-place ReLU for int8 buffer (`data[i] = max(data[i], 0)`).
31pub fn relu_s8(data: &mut [i8]) {
32    for val in data.iter_mut() {
33        if *val < 0 {
34            *val = 0;
35        }
36    }
37}
38
39/// In-place ReLU6 for int8 buffer (`data[i] = min(max(data[i], 0), 6)`).
40pub 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
53/// In-place generic activation clipping for int8 buffer (`data[i] = clamp(data[i], act.min, act.max)`).
54pub 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
61/// In-place ReLU for int16 buffer.
62pub fn relu_s16(data: &mut [i16]) {
63    for val in data.iter_mut() {
64        if *val < 0 {
65            *val = 0;
66        }
67    }
68}
69
70/// In-place generic activation clipping for int16 buffer.
71pub 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
78/// LeakyReLU activation for int8 buffer.
79pub 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
99/// Sigmoid activation for int8 tensors using direct lookup table.
100pub 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
115/// Tanh activation for int8 tensors using lookup table.
116pub 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
129/// Sigmoid activation for int16 tensors.
130pub 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
170/// Tanh activation for int16 tensors.
171pub 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        // Sigmoid(0) should be around 0 in s8 centered representation
232        assert!((output[0] as i32).abs() <= 5);
233        assert!(output[1] > 100);
234        assert!(output[2] < -100);
235    }
236}