1use crate::activations::{sigmoid_s16, tanh_s16};
4use crate::simd::{vec_dot_s16, vec_dot_s8};
5use crate::support::{clamp, requantize};
6use crate::types::{Activation, PerTensorQuantParams, Result};
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
10pub struct LstmGateParams {
11 pub input_offset: i32,
13 pub hidden_offset: i32,
15 pub multiplier: i32,
17 pub shift: i32,
19}
20
21pub fn lstm_step_s8_s16(
23 input: &[i8],
24 hidden_state: &mut [i8],
25 cell_state: &mut [i16],
26 weight_input: &[i8], weight_hidden: &[i8], bias: &[i32], gate_params: &LstmGateParams,
30 cell_clip: i16,
31 output_quant: &PerTensorQuantParams,
32 output_offset: i32,
33 activation: &Activation,
34) -> Result<()> {
35 let input_dim = input.len();
36 let hidden_dim = hidden_state.len();
37
38 for h in 0..hidden_dim {
39 let w_i_in = &weight_input[0 * hidden_dim * input_dim + h * input_dim..];
41 let w_f_in = &weight_input[1 * hidden_dim * input_dim + h * input_dim..];
42 let w_g_in = &weight_input[2 * hidden_dim * input_dim + h * input_dim..];
43 let w_o_in = &weight_input[3 * hidden_dim * input_dim + h * input_dim..];
44
45 let w_i_h = &weight_hidden[0 * hidden_dim * hidden_dim + h * hidden_dim..];
46 let w_f_h = &weight_hidden[1 * hidden_dim * hidden_dim + h * hidden_dim..];
47 let w_g_h = &weight_hidden[2 * hidden_dim * hidden_dim + h * hidden_dim..];
48 let w_o_h = &weight_hidden[3 * hidden_dim * hidden_dim + h * hidden_dim..];
49
50 let acc_i = bias[0 * hidden_dim + h]
52 + vec_dot_s8(input, w_i_in, gate_params.input_offset)
53 + vec_dot_s8(hidden_state, w_i_h, gate_params.hidden_offset);
54 let acc_f = bias[1 * hidden_dim + h]
55 + vec_dot_s8(input, w_f_in, gate_params.input_offset)
56 + vec_dot_s8(hidden_state, w_f_h, gate_params.hidden_offset);
57 let acc_g = bias[2 * hidden_dim + h]
58 + vec_dot_s8(input, w_g_in, gate_params.input_offset)
59 + vec_dot_s8(hidden_state, w_g_h, gate_params.hidden_offset);
60 let acc_o = bias[3 * hidden_dim + h]
61 + vec_dot_s8(input, w_o_in, gate_params.input_offset)
62 + vec_dot_s8(hidden_state, w_o_h, gate_params.hidden_offset);
63
64 let req_i = requantize(acc_i, gate_params.multiplier, gate_params.shift) as i16;
66 let req_f = requantize(acc_f, gate_params.multiplier, gate_params.shift) as i16;
67 let req_g = requantize(acc_g, gate_params.multiplier, gate_params.shift) as i16;
68 let req_o = requantize(acc_o, gate_params.multiplier, gate_params.shift) as i16;
69
70 let mut gate_i = [0i16];
71 let mut gate_f = [0i16];
72 let mut gate_g = [0i16];
73 let mut gate_o = [0i16];
74
75 sigmoid_s16(&[req_i], &mut gate_i, 0);
76 sigmoid_s16(&[req_f], &mut gate_f, 0);
77 tanh_s16(&[req_g], &mut gate_g, 0);
78 sigmoid_s16(&[req_o], &mut gate_o, 0);
79
80 let c_prev = cell_state[h] as i32;
82 let f_val = (gate_f[0] as i32 + 32768) >> 1; let i_val = (gate_i[0] as i32 + 32768) >> 1;
84 let g_val = gate_g[0] as i32;
85
86 let c_next = ((f_val * c_prev) >> 15) + ((i_val * g_val) >> 15);
87 let c_clamped = clamp(c_next, -cell_clip as i32, cell_clip as i32) as i16;
88 cell_state[h] = c_clamped;
89
90 let mut tan_c = [0i16];
92 tanh_s16(&[c_clamped], &mut tan_c, 0);
93
94 let o_val = (gate_o[0] as i32 + 32768) >> 1;
95 let h_next_raw = (o_val * (tan_c[0] as i32)) >> 15;
96 let h_req = requantize(h_next_raw, output_quant.multiplier, output_quant.shift);
97 let h_final = clamp(h_req + output_offset, activation.min, activation.max);
98
99 hidden_state[h] = h_final as i8;
100 }
101
102 Ok(())
103}
104
105pub fn lstm_step_s16(
107 input: &[i16],
108 hidden_state: &mut [i16],
109 cell_state: &mut [i16],
110 weight_input: &[i8],
111 weight_hidden: &[i8],
112 bias: &[i64],
113 gate_params: &LstmGateParams,
114 cell_clip: i16,
115 output_quant: &PerTensorQuantParams,
116 activation: &Activation,
117) -> Result<()> {
118 let input_dim = input.len();
119 let hidden_dim = hidden_state.len();
120
121 for h in 0..hidden_dim {
122 let w_i_in = &weight_input[0 * hidden_dim * input_dim + h * input_dim..];
123 let w_f_in = &weight_input[1 * hidden_dim * input_dim + h * input_dim..];
124 let w_g_in = &weight_input[2 * hidden_dim * input_dim + h * input_dim..];
125 let w_o_in = &weight_input[3 * hidden_dim * input_dim + h * input_dim..];
126
127 let w_i_h = &weight_hidden[0 * hidden_dim * hidden_dim + h * hidden_dim..];
128 let w_f_h = &weight_hidden[1 * hidden_dim * hidden_dim + h * hidden_dim..];
129 let w_g_h = &weight_hidden[2 * hidden_dim * hidden_dim + h * hidden_dim..];
130 let w_o_h = &weight_hidden[3 * hidden_dim * hidden_dim + h * hidden_dim..];
131
132 let mut acc_i = bias[0 * hidden_dim + h];
133 let mut acc_f = bias[1 * hidden_dim + h];
134 let mut acc_g = bias[2 * hidden_dim + h];
135 let mut acc_o = bias[3 * hidden_dim + h];
136
137 for i in 0..input_dim {
138 let in_val = input[i] as i64;
139 acc_i += in_val * (w_i_in[i] as i64);
140 acc_f += in_val * (w_f_in[i] as i64);
141 acc_g += in_val * (w_g_in[i] as i64);
142 acc_o += in_val * (w_o_in[i] as i64);
143 }
144
145 for i in 0..hidden_dim {
146 let h_val = hidden_state[i] as i64;
147 acc_i += h_val * (w_i_h[i] as i64);
148 acc_f += h_val * (w_f_h[i] as i64);
149 acc_g += h_val * (w_g_h[i] as i64);
150 acc_o += h_val * (w_o_h[i] as i64);
151 }
152
153 let req_i = requantize(
154 (acc_i >> 15) as i32,
155 gate_params.multiplier,
156 gate_params.shift,
157 ) as i16;
158 let req_f = requantize(
159 (acc_f >> 15) as i32,
160 gate_params.multiplier,
161 gate_params.shift,
162 ) as i16;
163 let req_g = requantize(
164 (acc_g >> 15) as i32,
165 gate_params.multiplier,
166 gate_params.shift,
167 ) as i16;
168 let req_o = requantize(
169 (acc_o >> 15) as i32,
170 gate_params.multiplier,
171 gate_params.shift,
172 ) as i16;
173
174 let mut gate_i = [0i16];
175 let mut gate_f = [0i16];
176 let mut gate_g = [0i16];
177 let mut gate_o = [0i16];
178
179 sigmoid_s16(&[req_i], &mut gate_i, 0);
180 sigmoid_s16(&[req_f], &mut gate_f, 0);
181 tanh_s16(&[req_g], &mut gate_g, 0);
182 sigmoid_s16(&[req_o], &mut gate_o, 0);
183
184 let c_prev = cell_state[h] as i32;
185 let f_val = (gate_f[0] as i32 + 32768) >> 1;
186 let i_val = (gate_i[0] as i32 + 32768) >> 1;
187 let g_val = gate_g[0] as i32;
188
189 let c_next = ((f_val * c_prev) >> 15) + ((i_val * g_val) >> 15);
190 let c_clamped = clamp(c_next, -cell_clip as i32, cell_clip as i32) as i16;
191 cell_state[h] = c_clamped;
192
193 let mut tan_c = [0i16];
194 tanh_s16(&[c_clamped], &mut tan_c, 0);
195
196 let o_val = (gate_o[0] as i32 + 32768) >> 1;
197 let h_next_raw = (o_val * (tan_c[0] as i32)) >> 15;
198 let h_req = requantize(h_next_raw, output_quant.multiplier, output_quant.shift);
199 let h_final = clamp(h_req, activation.min, activation.max);
200
201 hidden_state[h] = h_final as i16;
202 }
203
204 Ok(())
205}
206
207pub fn svdf_s8(
209 input_offset: i32,
210 output_offset: i32,
211 rank: usize,
212 input: &[i8],
213 state: &mut [i8],
214 weights_feature: &[i8],
215 weights_time: &[i8],
216 bias: Option<&[i32]>,
217 input_quant: &PerTensorQuantParams,
218 output_quant: &PerTensorQuantParams,
219 activation: &Activation,
220 output: &mut [i8],
221) -> Result<()> {
222 let input_dim = input.len();
223 let feature_dim = weights_feature.len() / input_dim;
224 let time_steps = weights_time.len() / feature_dim;
225 let units = feature_dim / rank;
226
227 for f in 0..feature_dim {
228 let state_slice = &mut state[f * time_steps..(f + 1) * time_steps];
229 state_slice.copy_within(1..time_steps, 0);
230 }
231
232 for f in 0..feature_dim {
233 let wf = &weights_feature[f * input_dim..(f + 1) * input_dim];
234 let acc = vec_dot_s8(input, wf, input_offset);
235 let req = requantize(acc, input_quant.multiplier, input_quant.shift);
236 let clamped = clamp(req, i8::MIN as i32, i8::MAX as i32);
237
238 state[f * time_steps + (time_steps - 1)] = clamped as i8;
239 }
240
241 for u in 0..units {
242 let mut acc = match bias {
243 Some(b) => b[u],
244 None => 0,
245 };
246
247 for r in 0..rank {
248 let f = u * rank + r;
249 let st = &state[f * time_steps..(f + 1) * time_steps];
250 let wt = &weights_time[f * time_steps..(f + 1) * time_steps];
251
252 acc += vec_dot_s8(st, wt, 0);
253 }
254
255 let req = requantize(acc, output_quant.multiplier, output_quant.shift);
256 let final_val = clamp(req + output_offset, activation.min, activation.max);
257 output[u] = final_val as i8;
258 }
259
260 Ok(())
261}
262
263pub fn svdf_state_s16_s8(
265 input_offset: i32,
266 output_offset: i32,
267 rank: usize,
268 input: &[i8],
269 state: &mut [i16], weights_feature: &[i8],
271 weights_time: &[i16],
272 bias: Option<&[i32]>,
273 input_quant: &PerTensorQuantParams,
274 output_quant: &PerTensorQuantParams,
275 activation: &Activation,
276 output: &mut [i8],
277) -> Result<()> {
278 let input_dim = input.len();
279 let feature_dim = weights_feature.len() / input_dim;
280 let time_steps = weights_time.len() / feature_dim;
281 let units = feature_dim / rank;
282
283 for f in 0..feature_dim {
284 let state_slice = &mut state[f * time_steps..(f + 1) * time_steps];
285 state_slice.copy_within(1..time_steps, 0);
286 }
287
288 for f in 0..feature_dim {
289 let wf = &weights_feature[f * input_dim..(f + 1) * input_dim];
290 let acc = vec_dot_s8(input, wf, input_offset);
291 let req = requantize(acc, input_quant.multiplier, input_quant.shift);
292 let clamped = clamp(req, i16::MIN as i32, i16::MAX as i32);
293
294 state[f * time_steps + (time_steps - 1)] = clamped as i16;
295 }
296
297 for u in 0..units {
298 let mut acc = match bias {
299 Some(b) => b[u] as i64,
300 None => 0i64,
301 };
302
303 for r in 0..rank {
304 let f = u * rank + r;
305 let st = &state[f * time_steps..(f + 1) * time_steps];
306 let wt = &weights_time[f * time_steps..(f + 1) * time_steps];
307
308 acc += vec_dot_s16(st, wt);
309 }
310
311 let req = requantize(
312 (acc >> 15) as i32,
313 output_quant.multiplier,
314 output_quant.shift,
315 );
316 let final_val = clamp(req + output_offset, activation.min, activation.max);
317 output[u] = final_val as i8;
318 }
319
320 Ok(())
321}
322
323#[cfg(test)]
324mod tests {
325 use super::*;
326
327 #[test]
328 fn test_svdf_s8_basic() {
329 let rank = 1usize;
330 let input = [10i8, 20i8];
331 let mut state = [0i8; 2];
332 let weights_feature = [1i8, 1i8];
333 let weights_time = [1i8, 1i8];
334 let bias = [0i32];
335
336 let input_quant = PerTensorQuantParams::new(1073741824, 0);
337 let output_quant = PerTensorQuantParams::new(1073741824, 0);
338 let act = Activation::int8_unconstrained();
339 let mut output = [0i8; 1];
340
341 svdf_s8(
342 0,
343 0,
344 rank,
345 &input,
346 &mut state,
347 &weights_feature,
348 &weights_time,
349 Some(&bias),
350 &input_quant,
351 &output_quant,
352 &act,
353 &mut output,
354 )
355 .unwrap();
356
357 assert_eq!(output[0], 8);
358 }
359}