Skip to main content

embedded_nn/
softmax.rs

1//! Softmax activation operations for quantized neural networks.
2
3use crate::support::{clamp, divide_by_power_of_two, doubling_high_mult_no_sat};
4use crate::types::Result;
5
6const ACCUM_BITS: i32 = 12;
7
8/// Saturation multiply (Q31 high multiplication).
9#[inline]
10fn mul_sat(a: i32, b: i32) -> i32 {
11    doubling_high_mult_no_sat(a, b)
12}
13
14/// Evaluates fixed-point exponent on negative values.
15pub fn exp_on_negative_values(val: i32) -> i32 {
16    if val == 0 {
17        return i32::MAX;
18    }
19
20    let mut shift = 24i32;
21    let val_mod_minus_quarter = (val & ((1 << shift) - 1)) - (1 << shift);
22    let remainder = val_mod_minus_quarter - val;
23    let x = (val_mod_minus_quarter << 5) + (1 << 28);
24    let x2 = mul_sat(x, x);
25
26    let op1 = divide_by_power_of_two(mul_sat(x2, x2), 2) + mul_sat(x2, x);
27    let op2 = x + divide_by_power_of_two(mul_sat(op1, 715827883) + x2, 1);
28    let mut result = 1895147668 + mul_sat(1895147668, op2);
29
30    let constants = [
31        1672461947, 1302514674, 790015084, 290630308, 39332535, 720401, 242,
32    ];
33
34    for &c in constants.iter() {
35        if (remainder & (1 << shift)) != 0 {
36            result = mul_sat(result, c);
37        }
38        shift += 1;
39    }
40
41    result
42}
43
44/// Evaluates 1 / (1 + x) for x in [0, 1] in fixed-point.
45pub fn one_over_one_plus_x_for_x_in_0_1(val: i32) -> i32 {
46    let sum = val as i64 + i32::MAX as i64;
47    let half_denominator = ((sum + if sum >= 0 { 1 } else { -1 }) / 2) as i32;
48    let mut x = 1515870810 + mul_sat(half_denominator, -1010580540);
49
50    let shift = 1i32 << 29;
51    for _ in 0..3 {
52        let diff = shift - mul_sat(half_denominator, x);
53        x += mul_sat(x, diff) << 2;
54    }
55
56    x << 1
57}
58
59/// Performs Softmax for int8 tensors.
60pub fn softmax_s8(
61    input: &[i8],
62    num_rows: usize,
63    row_size: usize,
64    mult: i32,
65    shift: i32,
66    diff_min: i32,
67    output: &mut [i8],
68) -> Result<()> {
69    let mask = 1i32 << shift;
70
71    for row in 0..num_rows {
72        let in_row = &input[row * row_size..(row + 1) * row_size];
73        let out_row = &mut output[row * row_size..(row + 1) * row_size];
74
75        // 1. Find max
76        let mut max_val = in_row[0];
77        for &val in &in_row[1..] {
78            if val > max_val {
79                max_val = val;
80            }
81        }
82
83        // 2. Accumulate sum of exp
84        let mut sum = 0i32;
85        for &val in in_row.iter() {
86            let diff = val as i32 - max_val as i32;
87            if diff >= diff_min {
88                let exp_input = mul_sat(diff * mask, mult);
89                let exp_val = exp_on_negative_values(exp_input);
90                sum += divide_by_power_of_two(exp_val, ACCUM_BITS);
91            }
92        }
93
94        // 3. Requantize output
95        let headroom = if sum > 0 {
96            sum.leading_zeros() as i32
97        } else {
98            32
99        };
100        let shifted_sum = if sum > 0 {
101            (sum << headroom) - (1i32 << 31)
102        } else {
103            0
104        };
105        let shifted_scale = one_over_one_plus_x_for_x_in_0_1(shifted_sum);
106        let bits_over_unit = ACCUM_BITS - headroom + 23;
107
108        for (i, &val) in in_row.iter().enumerate() {
109            let diff = val as i32 - max_val as i32;
110            if diff >= diff_min {
111                let exp_input = mul_sat(diff * mask, mult);
112                let exp_val = exp_on_negative_values(exp_input);
113                let scaled = mul_sat(shifted_scale, exp_val);
114                let res = divide_by_power_of_two(scaled, bits_over_unit) + (i8::MIN as i32);
115                out_row[i] = clamp(res, i8::MIN as i32, i8::MAX as i32) as i8;
116            } else {
117                out_row[i] = i8::MIN;
118            }
119        }
120    }
121    Ok(())
122}
123
124/// Performs Softmax for int16 tensors.
125pub fn softmax_s16(
126    input: &[i16],
127    num_rows: usize,
128    row_size: usize,
129    mult: i32,
130    shift: i32,
131    diff_min: i32,
132    output: &mut [i16],
133) -> Result<()> {
134    let mask = 1i32 << shift;
135
136    for row in 0..num_rows {
137        let in_row = &input[row * row_size..(row + 1) * row_size];
138        let out_row = &mut output[row * row_size..(row + 1) * row_size];
139
140        let mut max_val = in_row[0];
141        for &val in &in_row[1..] {
142            if val > max_val {
143                max_val = val;
144            }
145        }
146
147        let mut sum = 0i32;
148        for &val in in_row.iter() {
149            let diff = val as i32 - max_val as i32;
150            if diff >= diff_min {
151                let exp_input = mul_sat(diff * mask, mult);
152                let exp_val = exp_on_negative_values(exp_input);
153                sum += divide_by_power_of_two(exp_val, ACCUM_BITS);
154            }
155        }
156
157        let headroom = if sum > 0 {
158            sum.leading_zeros() as i32
159        } else {
160            32
161        };
162        let shifted_sum = if sum > 0 {
163            (sum << headroom) - (1i32 << 31)
164        } else {
165            0
166        };
167        let shifted_scale = one_over_one_plus_x_for_x_in_0_1(shifted_sum);
168        let bits_over_unit = ACCUM_BITS - headroom + 15;
169
170        for (i, &val) in in_row.iter().enumerate() {
171            let diff = val as i32 - max_val as i32;
172            if diff >= diff_min {
173                let exp_input = mul_sat(diff * mask, mult);
174                let exp_val = exp_on_negative_values(exp_input);
175                let scaled = mul_sat(shifted_scale, exp_val);
176                let res = divide_by_power_of_two(scaled, bits_over_unit) + (i16::MIN as i32);
177                out_row[i] = clamp(res, i16::MIN as i32, i16::MAX as i32) as i16;
178            } else {
179                out_row[i] = i16::MIN;
180            }
181        }
182    }
183    Ok(())
184}
185
186#[cfg(test)]
187mod tests {
188    use super::*;
189
190    #[test]
191    fn test_softmax_s8() {
192        let input = [10i8, 20i8, 30i8, 40i8];
193        let mut output = [0i8; 4];
194
195        softmax_s8(&input, 1, 4, 1073741824, 20, -256, &mut output).unwrap();
196
197        // Check that highest element input[3] produces maximum softmax probability
198        assert!(output[3] > output[2]);
199        assert!(output[2] > output[1]);
200        assert!(output[1] > output[0]);
201    }
202}