Skip to main content

embedded_nn/
convolution.rs

1//! Convolution layer operations for quantized neural networks.
2
3use crate::support::{clamp, requantize};
4use crate::types::{
5    ConvParams, Dims, DwConvParams, Error, PerChannelQuantParams, PerTensorQuantParams, Result,
6};
7
8/// Performs standard 2D Convolution with per-tensor quantization.
9pub fn convolve_s8(
10    conv_params: &ConvParams,
11    quant_params: &PerTensorQuantParams,
12    input_dims: &Dims,
13    input: &[i8],
14    filter_dims: &Dims,
15    kernel: &[i8],
16    bias: Option<&[i32]>,
17    output_dims: &Dims,
18    output: &mut [i8],
19) -> Result<()> {
20    let input_batches = input_dims.n as usize;
21    let input_h = input_dims.h as usize;
22    let input_w = input_dims.w as usize;
23    let input_c = input_dims.c as usize;
24
25    let kernel_h = filter_dims.h as usize;
26    let kernel_w = filter_dims.w as usize;
27    let kernel_c = filter_dims.c as usize;
28
29    let output_h = output_dims.h as usize;
30    let output_w = output_dims.w as usize;
31    let output_c = output_dims.c as usize;
32
33    if input_c == 0 || output_c == 0 {
34        return Err(Error::ArgumentError);
35    }
36
37    let groups = input_c / kernel_c;
38    let output_c_per_group = output_c / groups;
39
40    for b in 0..input_batches {
41        for out_y in 0..output_h {
42            let base_y = out_y as i32 * conv_params.stride.h - conv_params.padding.h;
43            for out_x in 0..output_w {
44                let base_x = out_x as i32 * conv_params.stride.w - conv_params.padding.w;
45
46                for g in 0..groups {
47                    for out_ch_idx in 0..output_c_per_group {
48                        let out_c = g * output_c_per_group + out_ch_idx;
49                        let mut acc: i32 = match bias {
50                            Some(b_slice) => b_slice[out_c],
51                            None => 0,
52                        };
53
54                        for ky in 0..kernel_h {
55                            let in_y = base_y + ky as i32 * conv_params.dilation.h;
56                            if in_y >= 0 && in_y < input_dims.h {
57                                for kx in 0..kernel_w {
58                                    let in_x = base_x + kx as i32 * conv_params.dilation.w;
59                                    if in_x >= 0 && in_x < input_dims.w {
60                                        let in_idx_base = ((b * input_h + in_y as usize) * input_w
61                                            + in_x as usize)
62                                            * input_c
63                                            + g * kernel_c;
64                                        let ker_idx_base =
65                                            ((out_c * kernel_h + ky) * kernel_w + kx) * kernel_c;
66
67                                        for ic in 0..kernel_c {
68                                            let lhs = input[in_idx_base + ic] as i32
69                                                + conv_params.input_offset;
70                                            let rhs = kernel[ker_idx_base + ic] as i32;
71                                            acc += lhs * rhs;
72                                        }
73                                    }
74                                }
75                            }
76                        }
77
78                        acc = requantize(acc, quant_params.multiplier, quant_params.shift);
79                        acc += conv_params.output_offset;
80                        acc = clamp(acc, conv_params.activation.min, conv_params.activation.max);
81
82                        let out_idx =
83                            ((b * output_h + out_y) * output_w + out_x) * output_c + out_c;
84                        output[out_idx] = acc as i8;
85                    }
86                }
87            }
88        }
89    }
90
91    Ok(())
92}
93
94/// Performs standard 2D Convolution with per-channel quantization.
95pub fn convolve_per_channel_s8(
96    conv_params: &ConvParams,
97    quant_params: &PerChannelQuantParams,
98    input_dims: &Dims,
99    input: &[i8],
100    filter_dims: &Dims,
101    kernel: &[i8],
102    bias: Option<&[i32]>,
103    output_dims: &Dims,
104    output: &mut [i8],
105) -> Result<()> {
106    let input_batches = input_dims.n as usize;
107    let input_h = input_dims.h as usize;
108    let input_w = input_dims.w as usize;
109    let input_c = input_dims.c as usize;
110
111    let kernel_h = filter_dims.h as usize;
112    let kernel_w = filter_dims.w as usize;
113    let kernel_c = filter_dims.c as usize;
114
115    let output_h = output_dims.h as usize;
116    let output_w = output_dims.w as usize;
117    let output_c = output_dims.c as usize;
118
119    if input_c == 0 || output_c == 0 {
120        return Err(Error::ArgumentError);
121    }
122
123    let groups = input_c / kernel_c;
124    let output_c_per_group = output_c / groups;
125
126    for b in 0..input_batches {
127        for out_y in 0..output_h {
128            let base_y = out_y as i32 * conv_params.stride.h - conv_params.padding.h;
129            for out_x in 0..output_w {
130                let base_x = out_x as i32 * conv_params.stride.w - conv_params.padding.w;
131
132                for g in 0..groups {
133                    for out_ch_idx in 0..output_c_per_group {
134                        let out_c = g * output_c_per_group + out_ch_idx;
135                        let mut acc: i32 = match bias {
136                            Some(b_slice) => b_slice[out_c],
137                            None => 0,
138                        };
139
140                        for ky in 0..kernel_h {
141                            let in_y = base_y + ky as i32 * conv_params.dilation.h;
142                            if in_y >= 0 && in_y < input_dims.h {
143                                for kx in 0..kernel_w {
144                                    let in_x = base_x + kx as i32 * conv_params.dilation.w;
145                                    if in_x >= 0 && in_x < input_dims.w {
146                                        let in_idx_base = ((b * input_h + in_y as usize) * input_w
147                                            + in_x as usize)
148                                            * input_c
149                                            + g * kernel_c;
150                                        let ker_idx_base =
151                                            ((out_c * kernel_h + ky) * kernel_w + kx) * kernel_c;
152
153                                        for ic in 0..kernel_c {
154                                            let lhs = input[in_idx_base + ic] as i32
155                                                + conv_params.input_offset;
156                                            let rhs = kernel[ker_idx_base + ic] as i32;
157                                            acc += lhs * rhs;
158                                        }
159                                    }
160                                }
161                            }
162                        }
163
164                        let mult = quant_params.multiplier[out_c];
165                        let shift = quant_params.shift[out_c];
166
167                        acc = requantize(acc, mult, shift);
168                        acc += conv_params.output_offset;
169                        acc = clamp(acc, conv_params.activation.min, conv_params.activation.max);
170
171                        let out_idx =
172                            ((b * output_h + out_y) * output_w + out_x) * output_c + out_c;
173                        output[out_idx] = acc as i8;
174                    }
175                }
176            }
177        }
178    }
179
180    Ok(())
181}
182
183/// Performs Depthwise 2D Convolution with per-channel quantization.
184pub fn depthwise_conv_per_channel_s8(
185    dw_params: &DwConvParams,
186    quant_params: &PerChannelQuantParams,
187    input_dims: &Dims,
188    input: &[i8],
189    filter_dims: &Dims,
190    kernel: &[i8],
191    bias: Option<&[i32]>,
192    output_dims: &Dims,
193    output: &mut [i8],
194) -> Result<()> {
195    let input_batches = input_dims.n as usize;
196    let input_h = input_dims.h as usize;
197    let input_w = input_dims.w as usize;
198    let input_c = input_dims.c as usize;
199
200    let kernel_h = filter_dims.h as usize;
201    let kernel_w = filter_dims.w as usize;
202
203    let output_h = output_dims.h as usize;
204    let output_w = output_dims.w as usize;
205    let output_c = output_dims.c as usize;
206
207    let ch_mult = dw_params.ch_mult as usize;
208
209    for b in 0..input_batches {
210        for out_y in 0..output_h {
211            let base_y = out_y as i32 * dw_params.stride.h - dw_params.padding.h;
212            for out_x in 0..output_w {
213                let base_x = out_x as i32 * dw_params.stride.w - dw_params.padding.w;
214
215                for out_c in 0..output_c {
216                    let in_c = out_c / ch_mult;
217
218                    let mut acc: i32 = match bias {
219                        Some(b_slice) => b_slice[out_c],
220                        None => 0,
221                    };
222
223                    for ky in 0..kernel_h {
224                        let in_y = base_y + ky as i32 * dw_params.dilation.h;
225                        if in_y >= 0 && in_y < input_dims.h {
226                            for kx in 0..kernel_w {
227                                let in_x = base_x + kx as i32 * dw_params.dilation.w;
228                                if in_x >= 0 && in_x < input_dims.w {
229                                    let in_idx = ((b * input_h + in_y as usize) * input_w
230                                        + in_x as usize)
231                                        * input_c
232                                        + in_c;
233                                    let ker_idx = ((ky * kernel_w + kx) * output_c) + out_c;
234
235                                    let lhs = input[in_idx] as i32 + dw_params.input_offset;
236                                    let rhs = kernel[ker_idx] as i32;
237                                    acc += lhs * rhs;
238                                }
239                            }
240                        }
241                    }
242
243                    let mult = quant_params.multiplier[out_c];
244                    let shift = quant_params.shift[out_c];
245
246                    acc = requantize(acc, mult, shift);
247                    acc += dw_params.output_offset;
248                    acc = clamp(acc, dw_params.activation.min, dw_params.activation.max);
249
250                    let out_idx = ((b * output_h + out_y) * output_w + out_x) * output_c + out_c;
251                    output[out_idx] = acc as i8;
252                }
253            }
254        }
255    }
256
257    Ok(())
258}
259
260/// Performs Transposed 2D Convolution (Deconvolution) for int8 tensors with per-channel quantization.
261pub fn transpose_conv_s8(
262    conv_params: &ConvParams,
263    quant_params: &PerChannelQuantParams,
264    input_dims: &Dims,
265    input: &[i8],
266    filter_dims: &Dims,
267    kernel: &[i8],
268    bias: Option<&[i32]>,
269    output_dims: &Dims,
270    output: &mut [i8],
271) -> Result<()> {
272    let input_batches = input_dims.n as usize;
273    let input_h = input_dims.h as usize;
274    let input_w = input_dims.w as usize;
275    let input_c = input_dims.c as usize;
276
277    let kernel_h = filter_dims.h as usize;
278    let kernel_w = filter_dims.w as usize;
279
280    let output_h = output_dims.h as usize;
281    let output_w = output_dims.w as usize;
282    let output_c = output_dims.c as usize;
283
284    let stride_y = conv_params.stride.h as usize;
285    let stride_x = conv_params.stride.w as usize;
286    let pad_y = conv_params.padding.h as usize;
287    let pad_x = conv_params.padding.w as usize;
288
289    let mut accum_buffer = [0i32; 1024]; // Stack scratch accumulator for 1024 elements chunk or dynamic iteration
290    let scratch_size = output_h * output_w * output_c;
291
292    for b in 0..input_batches {
293        for out_idx in 0..scratch_size {
294            let out_c = out_idx % output_c;
295            let b_val = match bias {
296                Some(b_slice) => b_slice[out_c],
297                None => 0,
298            };
299            if out_idx < accum_buffer.len() {
300                accum_buffer[out_idx] = b_val;
301            }
302        }
303
304        // Scatter-accumulate input into output positions
305        for in_y in 0..input_h {
306            for in_x in 0..input_w {
307                for ky in 0..kernel_h {
308                    let out_y = in_y * stride_y + ky;
309                    if out_y >= pad_y && out_y < output_h + pad_y {
310                        let actual_out_y = out_y - pad_y;
311                        for kx in 0..kernel_w {
312                            let out_x = in_x * stride_x + kx;
313                            if out_x >= pad_x && out_x < output_w + pad_x {
314                                let actual_out_x = out_x - pad_x;
315
316                                for out_c in 0..output_c {
317                                    for in_c in 0..input_c {
318                                        let in_idx = ((b * input_h + in_y) * input_w + in_x)
319                                            * input_c
320                                            + in_c;
321                                        let ker_idx = ((out_c * kernel_h + ky) * kernel_w + kx)
322                                            * input_c
323                                            + in_c;
324
325                                        let lhs = input[in_idx] as i32 + conv_params.input_offset;
326                                        let rhs = kernel[ker_idx] as i32;
327
328                                        let buf_idx = (actual_out_y * output_w + actual_out_x)
329                                            * output_c
330                                            + out_c;
331                                        if buf_idx < accum_buffer.len() {
332                                            accum_buffer[buf_idx] += lhs * rhs;
333                                        }
334                                    }
335                                }
336                            }
337                        }
338                    }
339                }
340            }
341        }
342
343        // Requantize and write back output
344        for out_y in 0..output_h {
345            for out_x in 0..output_w {
346                for out_c in 0..output_c {
347                    let buf_idx = (out_y * output_w + out_x) * output_c + out_c;
348                    let acc = if buf_idx < accum_buffer.len() {
349                        accum_buffer[buf_idx]
350                    } else {
351                        match bias {
352                            Some(b_slice) => b_slice[out_c],
353                            None => 0,
354                        }
355                    };
356
357                    let mult = quant_params.multiplier[out_c];
358                    let shift = quant_params.shift[out_c];
359
360                    let req = requantize(acc, mult, shift);
361                    let final_val = clamp(
362                        req + conv_params.output_offset,
363                        conv_params.activation.min,
364                        conv_params.activation.max,
365                    );
366
367                    let out_idx = ((b * output_h + out_y) * output_w + out_x) * output_c + out_c;
368                    if out_idx < output.len() {
369                        output[out_idx] = final_val as i8;
370                    }
371                }
372            }
373        }
374    }
375
376    Ok(())
377}
378
379/// Performs 1D Temporal Convolution for int8 tensors (`convolve_1_x_n_s8`).
380pub fn convolve_1_x_n_s8(
381    conv_params: &ConvParams,
382    quant_params: &PerTensorQuantParams,
383    input_dims: &Dims,
384    input: &[i8],
385    filter_dims: &Dims,
386    kernel: &[i8],
387    bias: Option<&[i32]>,
388    output_dims: &Dims,
389    output: &mut [i8],
390) -> Result<()> {
391    // 1D Conv is equivalent to 2D Conv with Height = 1
392    convolve_s8(
393        conv_params,
394        quant_params,
395        input_dims,
396        input,
397        filter_dims,
398        kernel,
399        bias,
400        output_dims,
401        output,
402    )
403}
404
405#[cfg(test)]
406mod tests {
407    use super::*;
408    use crate::types::{Activation, Tile};
409
410    #[test]
411    fn test_convolve_s8_simple() {
412        let conv_params = ConvParams {
413            input_offset: 0,
414            output_offset: 0,
415            stride: Tile::new(1, 1),
416            padding: Tile::new(0, 0),
417            dilation: Tile::new(1, 1),
418            activation: Activation::int8_unconstrained(),
419        };
420
421        let quant_params = PerTensorQuantParams::new(1073741824, 0); // 0.5
422
423        let input_dims = Dims::new(1, 3, 3, 1);
424        let input = [1i8, 2i8, 3i8, 4i8, 5i8, 6i8, 7i8, 8i8, 9i8];
425
426        let filter_dims = Dims::new(1, 2, 2, 1); // 1 out_channel, 2x2 kernel, 1 in_channel
427        let kernel = [1i8, 0i8, 0i8, 1i8];
428
429        let output_dims = Dims::new(1, 2, 2, 1);
430        let mut output = [0i8; 4];
431
432        convolve_s8(
433            &conv_params,
434            &quant_params,
435            &input_dims,
436            &input,
437            &filter_dims,
438            &kernel,
439            None,
440            &output_dims,
441            &mut output,
442        )
443        .unwrap();
444
445        assert_eq!(output, [3, 4, 6, 7]);
446    }
447}