Skip to main content

embedded_nn/
fully_connected.rs

1//! Fully Connected (Linear / Dense) and Batch Matrix Multiplication operations.
2
3use crate::support::{clamp, requantize};
4use crate::types::{Dims, FcParams, PerChannelQuantParams, PerTensorQuantParams, Result};
5
6/// Performs per-tensor quantized int8 Fully Connected layer.
7pub fn fully_connected_s8(
8    fc_params: &FcParams,
9    quant_params: &PerTensorQuantParams,
10    input_dims: &Dims,
11    input: &[i8],
12    filter_dims: &Dims,
13    kernel: &[i8],
14    bias: Option<&[i32]>,
15    output_dims: &Dims,
16    output: &mut [i8],
17) -> Result<()> {
18    let batches = input_dims.n as usize;
19    let accum_depth = filter_dims.n as usize;
20    let output_depth = output_dims.c as usize;
21
22    for b in 0..batches {
23        let input_batch = &input[b * accum_depth..(b + 1) * accum_depth];
24        let output_batch = &mut output[b * output_depth..(b + 1) * output_depth];
25
26        for out_c in 0..output_depth {
27            let mut acc: i32 = match bias {
28                Some(b_slice) => b_slice[out_c],
29                None => 0,
30            };
31
32            let kernel_row = &kernel[out_c * accum_depth..(out_c + 1) * accum_depth];
33
34            for i in 0..accum_depth {
35                let lhs = input_batch[i] as i32 + fc_params.input_offset;
36                let rhs = kernel_row[i] as i32 + fc_params.filter_offset;
37                acc += lhs * rhs;
38            }
39
40            acc = requantize(acc, quant_params.multiplier, quant_params.shift);
41            acc += fc_params.output_offset;
42            acc = clamp(acc, fc_params.activation.min, fc_params.activation.max);
43
44            output_batch[out_c] = acc as i8;
45        }
46    }
47    Ok(())
48}
49
50/// Performs per-channel quantized int8 Fully Connected layer.
51pub fn fully_connected_per_channel_s8(
52    fc_params: &FcParams,
53    quant_params: &PerChannelQuantParams,
54    input_dims: &Dims,
55    input: &[i8],
56    filter_dims: &Dims,
57    kernel: &[i8],
58    bias: Option<&[i32]>,
59    output_dims: &Dims,
60    output: &mut [i8],
61) -> Result<()> {
62    let batches = input_dims.n as usize;
63    let accum_depth = filter_dims.n as usize;
64    let output_depth = output_dims.c as usize;
65
66    for b in 0..batches {
67        let input_batch = &input[b * accum_depth..(b + 1) * accum_depth];
68        let output_batch = &mut output[b * output_depth..(b + 1) * output_depth];
69
70        for out_c in 0..output_depth {
71            let mut acc: i32 = match bias {
72                Some(b_slice) => b_slice[out_c],
73                None => 0,
74            };
75
76            let kernel_row = &kernel[out_c * accum_depth..(out_c + 1) * accum_depth];
77
78            for i in 0..accum_depth {
79                let lhs = input_batch[i] as i32 + fc_params.input_offset;
80                let rhs = kernel_row[i] as i32 + fc_params.filter_offset;
81                acc += lhs * rhs;
82            }
83
84            let mult = quant_params.multiplier[out_c];
85            let shift = quant_params.shift[out_c];
86
87            acc = requantize(acc, mult, shift);
88            acc += fc_params.output_offset;
89            acc = clamp(acc, fc_params.activation.min, fc_params.activation.max);
90
91            output_batch[out_c] = acc as i8;
92        }
93    }
94    Ok(())
95}
96
97/// Performs int16 Fully Connected layer.
98pub fn fully_connected_s16(
99    fc_params: &FcParams,
100    quant_params: &PerTensorQuantParams,
101    input_dims: &Dims,
102    input: &[i16],
103    filter_dims: &Dims,
104    kernel: &[i8],
105    bias: Option<&[i64]>,
106    output_dims: &Dims,
107    output: &mut [i16],
108) -> Result<()> {
109    let batches = input_dims.n as usize;
110    let accum_depth = filter_dims.n as usize;
111    let output_depth = output_dims.c as usize;
112
113    for b in 0..batches {
114        let input_batch = &input[b * accum_depth..(b + 1) * accum_depth];
115        let output_batch = &mut output[b * output_depth..(b + 1) * output_depth];
116
117        for out_c in 0..output_depth {
118            let mut acc: i64 = match bias {
119                Some(b_slice) => b_slice[out_c],
120                None => 0,
121            };
122
123            let kernel_row = &kernel[out_c * accum_depth..(out_c + 1) * accum_depth];
124
125            for i in 0..accum_depth {
126                let lhs = input_batch[i] as i64;
127                let rhs = kernel_row[i] as i64;
128                acc += lhs * rhs;
129            }
130
131            let req = requantize(
132                (acc >> 15) as i32,
133                quant_params.multiplier,
134                quant_params.shift,
135            );
136            let final_val = clamp(req, fc_params.activation.min, fc_params.activation.max);
137
138            output_batch[out_c] = final_val as i16;
139        }
140    }
141    Ok(())
142}
143
144/// Performs Batch Matrix Multiplication (`BatchMatMul`) for int8 tensors.
145///
146/// Computes `Output[b, i, j] = Requantize(sum_k (LHS[b, i, k] + lhs_offset) * (RHS[b, k, j] + rhs_offset))`
147pub fn batch_matmul_s8(
148    fc_params: &FcParams,
149    quant_params: &PerTensorQuantParams,
150    lhs_dims: &Dims,
151    input_lhs: &[i8],
152    rhs_dims: &Dims,
153    input_rhs: &[i8],
154    output_dims: &Dims,
155    output: &mut [i8],
156) -> Result<()> {
157    let batches = output_dims.n as usize;
158    let rows = lhs_dims.h as usize;
159    let cols = rhs_dims.c as usize;
160    let accum_dim = lhs_dims.w as usize;
161
162    for b in 0..batches {
163        let lhs_b_idx = if lhs_dims.n == 1 { 0 } else { b };
164        let rhs_b_idx = if rhs_dims.n == 1 { 0 } else { b };
165
166        for i in 0..rows {
167            for j in 0..cols {
168                let mut acc: i32 = 0;
169
170                for k in 0..accum_dim {
171                    let lhs_idx = (lhs_b_idx * rows + i) * accum_dim + k;
172                    let rhs_idx = (rhs_b_idx * accum_dim + k) * cols + j;
173
174                    let lhs_val = input_lhs[lhs_idx] as i32 + fc_params.input_offset;
175                    let rhs_val = input_rhs[rhs_idx] as i32 + fc_params.filter_offset;
176                    acc += lhs_val * rhs_val;
177                }
178
179                let req = requantize(acc, quant_params.multiplier, quant_params.shift);
180                let final_val = clamp(
181                    req + fc_params.output_offset,
182                    fc_params.activation.min,
183                    fc_params.activation.max,
184                );
185
186                let out_idx = (b * rows + i) * cols + j;
187                if out_idx < output.len() {
188                    output[out_idx] = final_val as i8;
189                }
190            }
191        }
192    }
193
194    Ok(())
195}
196
197/// Performs Batch Matrix Multiplication (`BatchMatMul`) for int16 tensors.
198pub fn batch_matmul_s16(
199    fc_params: &FcParams,
200    quant_params: &PerTensorQuantParams,
201    lhs_dims: &Dims,
202    input_lhs: &[i16],
203    rhs_dims: &Dims,
204    input_rhs: &[i16],
205    output_dims: &Dims,
206    output: &mut [i16],
207) -> Result<()> {
208    let batches = output_dims.n as usize;
209    let rows = lhs_dims.h as usize;
210    let cols = rhs_dims.c as usize;
211    let accum_dim = lhs_dims.w as usize;
212
213    for b in 0..batches {
214        let lhs_b_idx = if lhs_dims.n == 1 { 0 } else { b };
215        let rhs_b_idx = if rhs_dims.n == 1 { 0 } else { b };
216
217        for i in 0..rows {
218            for j in 0..cols {
219                let mut acc: i64 = 0;
220
221                for k in 0..accum_dim {
222                    let lhs_idx = (lhs_b_idx * rows + i) * accum_dim + k;
223                    let rhs_idx = (rhs_b_idx * accum_dim + k) * cols + j;
224
225                    let lhs_val = input_lhs[lhs_idx] as i64;
226                    let rhs_val = input_rhs[rhs_idx] as i64;
227                    acc += lhs_val * rhs_val;
228                }
229
230                let req = requantize(
231                    (acc >> 15) as i32,
232                    quant_params.multiplier,
233                    quant_params.shift,
234                );
235                let final_val = clamp(req, fc_params.activation.min, fc_params.activation.max);
236
237                let out_idx = (b * rows + i) * cols + j;
238                if out_idx < output.len() {
239                    output[out_idx] = final_val as i16;
240                }
241            }
242        }
243    }
244
245    Ok(())
246}
247
248#[cfg(test)]
249mod tests {
250    use super::*;
251    use crate::types::Activation;
252
253    #[test]
254    fn test_fully_connected_s8() {
255        let fc_params = FcParams {
256            input_offset: 0,
257            filter_offset: 0,
258            output_offset: 0,
259            activation: Activation::int8_unconstrained(),
260        };
261        let quant_params = PerTensorQuantParams::new(1073741824, 0); // 0.5
262        let input_dims = Dims::new(1, 1, 1, 3);
263        let input = [2i8, 4i8, 6i8];
264        let filter_dims = Dims::new(3, 1, 1, 2); // 3 input, 2 output channels
265        let kernel = [
266            1i8, 2i8, 3i8, // row 0
267            4i8, 5i8, 6i8, // row 1
268        ];
269        let bias = [0i32, 0i32];
270        let output_dims = Dims::new(1, 1, 1, 2);
271        let mut output = [0i8; 2];
272
273        fully_connected_s8(
274            &fc_params,
275            &quant_params,
276            &input_dims,
277            &input,
278            &filter_dims,
279            &kernel,
280            Some(&bias),
281            &output_dims,
282            &mut output,
283        )
284        .unwrap();
285
286        assert_eq!(output[0], 14);
287        assert_eq!(output[1], 32);
288    }
289
290    #[test]
291    fn test_batch_matmul_s8() {
292        let fc_params = FcParams {
293            input_offset: 0,
294            filter_offset: 0,
295            output_offset: 0,
296            activation: Activation::int8_unconstrained(),
297        };
298        let quant_params = PerTensorQuantParams::new(1073741824, 0); // 0.5
299
300        let lhs_dims = Dims::new(1, 2, 2, 0); // 2 rows x 2 cols
301        let input_lhs = [1i8, 2i8, 3i8, 4i8];
302
303        let rhs_dims = Dims::new(1, 2, 0, 2); // 2 rows x 2 cols
304        let input_rhs = [5i8, 6i8, 7i8, 8i8];
305
306        let output_dims = Dims::new(1, 2, 0, 2);
307        let mut output = [0i8; 4];
308
309        batch_matmul_s8(
310            &fc_params,
311            &quant_params,
312            &lhs_dims,
313            &input_lhs,
314            &rhs_dims,
315            &input_rhs,
316            &output_dims,
317            &mut output,
318        )
319        .unwrap();
320
321        // Row 0, Col 0: 1*5 + 2*7 = 19 -> *0.5 = 10 (rounded)
322        // Row 0, Col 1: 1*6 + 2*8 = 22 -> *0.5 = 11
323        // Row 1, Col 0: 3*5 + 4*7 = 43 -> *0.5 = 22
324        // Row 1, Col 1: 3*6 + 4*8 = 50 -> *0.5 = 25
325        assert_eq!(output, [10, 11, 22, 25]);
326    }
327}