Skip to main content

embedded_nn/
recurrent.rs

1//! Advanced recurrent neural network layers (LSTM, SVDF).
2
3use 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/// Quantized parameters for an LSTM cell gate.
9#[derive(Debug, Clone, Copy, PartialEq, Eq)]
10pub struct LstmGateParams {
11    /// Zero point offset for input.
12    pub input_offset: i32,
13    /// Zero point offset for recurrent hidden state.
14    pub hidden_offset: i32,
15    /// Requantization multiplier for gate pre-activation.
16    pub multiplier: i32,
17    /// Requantization shift for gate pre-activation.
18    pub shift: i32,
19}
20
21/// Unidirectional LSTM cell step for quantized int8/int16 state tensors (`lstm_step_s8_s16`).
22pub fn lstm_step_s8_s16(
23    input: &[i8],
24    hidden_state: &mut [i8],
25    cell_state: &mut [i16],
26    weight_input: &[i8],  // 4 * hidden_dim x input_dim (i, f, g, o)
27    weight_hidden: &[i8], // 4 * hidden_dim x hidden_dim (i, f, g, o)
28    bias: &[i32],         // 4 * hidden_dim
29    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        // Gates layout: 0: input (i), 1: forget (f), 2: cell (g), 3: output (o)
40        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        // 1. Accumulate pre-activations
51        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        // 2. Requantize pre-activations to int16 range
65        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        // 3. Compute cell state c_t = f * c_prev + i * g
81        let c_prev = cell_state[h] as i32;
82        let f_val = (gate_f[0] as i32 + 32768) >> 1; // Q15 scale
83        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        // 4. Compute hidden state h_t = o * tanh(c_t)
91        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
105/// Full int16 Unidirectional LSTM cell step (`lstm_step_s16`).
106pub 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
207/// SVDF (Singular Value Decomposition Filter) layer step for int8 tensors (`svdf_s8`).
208pub 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
263/// SVDF layer step with int16 state tensor (`svdf_state_s16_s8`).
264pub fn svdf_state_s16_s8(
265    input_offset: i32,
266    output_offset: i32,
267    rank: usize,
268    input: &[i8],
269    state: &mut [i16], // 16-bit state tensor for high precision
270    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}