1use crate::types::{Dims, Result, Tile};
4
5pub fn pad_s8(
7 input_dims: &Dims,
8 input: &[i8],
9 padding_before: &Tile,
10 _padding_after: &Tile,
11 pad_value: i8,
12 output_dims: &Dims,
13 output: &mut [i8],
14) -> Result<()> {
15 let batches = input_dims.n as usize;
16 let input_h = input_dims.h as usize;
17 let input_w = input_dims.w as usize;
18 let channels = input_dims.c as usize;
19
20 let pad_h_before = padding_before.h as usize;
21 let pad_w_before = padding_before.w as usize;
22
23 let output_h = output_dims.h as usize;
24 let output_w = output_dims.w as usize;
25
26 output.fill(pad_value);
27
28 for b in 0..batches {
29 for y in 0..input_h {
30 let out_y = y + pad_h_before;
31 if out_y < output_h {
32 for x in 0..input_w {
33 let out_x = x + pad_w_before;
34 if out_x < output_w {
35 let in_idx = ((b * input_h + y) * input_w + x) * channels;
36 let out_idx = ((b * output_h + out_y) * output_w + out_x) * channels;
37
38 output[out_idx..out_idx + channels]
39 .copy_from_slice(&input[in_idx..in_idx + channels]);
40 }
41 }
42 }
43 }
44 }
45
46 Ok(())
47}
48
49#[cfg(test)]
50mod tests {
51 use super::*;
52
53 #[test]
54 fn test_pad_s8() {
55 let in_dims = Dims::new(1, 2, 2, 1);
56 let input = [1i8, 2i8, 3i8, 4i8];
57
58 let pad_before = Tile::new(1, 1);
59 let pad_after = Tile::new(1, 1);
60
61 let out_dims = Dims::new(1, 4, 4, 1);
62 let mut out = [0i8; 16];
63
64 pad_s8(
65 &in_dims,
66 &input,
67 &pad_before,
68 &pad_after,
69 0i8,
70 &out_dims,
71 &mut out,
72 )
73 .unwrap();
74
75 assert_eq!(out[5], 1);
77 assert_eq!(out[6], 2);
78 assert_eq!(out[9], 3);
79 assert_eq!(out[10], 4);
80 assert_eq!(out[0], 0);
82 }
83}