1use crate::types::{Dims, Error, Result};
4
5pub 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}