Skip to main content

embedded_nn/
pad.rs

1//! Tensor padding operations for quantized tensors.
2
3use crate::types::{Dims, Result, Tile};
4
5/// Pads an int8 4D tensor with a specified pad value.
6pub 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        // Check center 2x2 elements match input
76        assert_eq!(out[5], 1);
77        assert_eq!(out[6], 2);
78        assert_eq!(out[9], 3);
79        assert_eq!(out[10], 4);
80        // Border element is 0
81        assert_eq!(out[0], 0);
82    }
83}