1use 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#[inline]
10fn mul_sat(a: i32, b: i32) -> i32 {
11 doubling_high_mult_no_sat(a, b)
12}
13
14pub 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
44pub 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
59pub 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 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 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 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
124pub 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 assert!(output[3] > output[2]);
199 assert!(output[2] > output[1]);
200 assert!(output[1] > output[0]);
201 }
202}