Skip to main content

embedded_nn/
concat.rs

1//! Concatenation operations along channel/depth dimension for quantized tensors.
2
3use crate::types::{Dims, Error, Result};
4
5/// Concatenates two int8 tensors along the channel (depth) dimension.
6pub fn concatenation_s8(
7    input1_dims: &Dims,
8    input1: &[i8],
9    input2_dims: &Dims,
10    input2: &[i8],
11    output_dims: &Dims,
12    output: &mut [i8],
13) -> Result<()> {
14    if input1_dims.n != input2_dims.n
15        || input1_dims.h != input2_dims.h
16        || input1_dims.w != input2_dims.w
17    {
18        return Err(Error::ArgumentError);
19    }
20
21    let out_c = input1_dims.c + input2_dims.c;
22    if output_dims.c != out_c {
23        return Err(Error::ArgumentError);
24    }
25
26    let outer_size = (input1_dims.n * input1_dims.h * input1_dims.w) as usize;
27    let c1 = input1_dims.c as usize;
28    let c2 = input2_dims.c as usize;
29
30    for i in 0..outer_size {
31        let in1_slice = &input1[i * c1..(i + 1) * c1];
32        let in2_slice = &input2[i * c2..(i + 1) * c2];
33        let out_slice = &mut output[i * (c1 + c2)..(i + 1) * (c1 + c2)];
34
35        out_slice[..c1].copy_from_slice(in1_slice);
36        out_slice[c1..c1 + c2].copy_from_slice(in2_slice);
37    }
38
39    Ok(())
40}
41
42#[cfg(test)]
43mod tests {
44    use super::*;
45
46    #[test]
47    fn test_concatenation_s8() {
48        let in1_dims = Dims::new(1, 1, 1, 2);
49        let in1 = [1i8, 2i8];
50        let in2_dims = Dims::new(1, 1, 1, 3);
51        let in2 = [3i8, 4i8, 5i8];
52
53        let out_dims = Dims::new(1, 1, 1, 5);
54        let mut out = [0i8; 5];
55
56        concatenation_s8(&in1_dims, &in1, &in2_dims, &in2, &out_dims, &mut out).unwrap();
57        assert_eq!(out, [1, 2, 3, 4, 5]);
58    }
59}