Skip to main content

embedded_nn/
pooling.rs

1//! Pooling layer operations (Max Pooling, Average Pooling) for quantized neural networks.
2
3use crate::support::clamp;
4use crate::types::{Dims, PoolParams, Result, Tile};
5
6/// Performs Max Pooling 2D for int8 tensors.
7pub fn max_pool_s8(
8    pool_params: &PoolParams,
9    filter_dims: &Tile,
10    input_dims: &Dims,
11    input: &[i8],
12    output_dims: &Dims,
13    output: &mut [i8],
14) -> Result<()> {
15    let input_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 kernel_h = filter_dims.h as usize;
21    let kernel_w = filter_dims.w as usize;
22
23    let output_h = output_dims.h as usize;
24    let output_w = output_dims.w as usize;
25
26    for b in 0..input_batches {
27        for out_y in 0..output_h {
28            let base_y = out_y as i32 * pool_params.stride.h - pool_params.padding.h;
29            for out_x in 0..output_w {
30                let base_x = out_x as i32 * pool_params.stride.w - pool_params.padding.w;
31
32                for c in 0..channels {
33                    let mut max_val = i8::MIN as i32;
34
35                    for ky in 0..kernel_h {
36                        let in_y = base_y + ky as i32;
37                        if in_y >= 0 && in_y < input_dims.h {
38                            for kx in 0..kernel_w {
39                                let in_x = base_x + kx as i32;
40                                if in_x >= 0 && in_x < input_dims.w {
41                                    let in_idx = ((b * input_h + in_y as usize) * input_w
42                                        + in_x as usize)
43                                        * channels
44                                        + c;
45                                    let val = input[in_idx] as i32;
46                                    if val > max_val {
47                                        max_val = val;
48                                    }
49                                }
50                            }
51                        }
52                    }
53
54                    max_val = clamp(
55                        max_val,
56                        pool_params.activation.min,
57                        pool_params.activation.max,
58                    );
59                    let out_idx = ((b * output_h + out_y) * output_w + out_x) * channels + c;
60                    output[out_idx] = max_val as i8;
61                }
62            }
63        }
64    }
65    Ok(())
66}
67
68/// Performs Average Pooling 2D for int8 tensors.
69pub fn avg_pool_s8(
70    pool_params: &PoolParams,
71    filter_dims: &Tile,
72    input_dims: &Dims,
73    input: &[i8],
74    output_dims: &Dims,
75    output: &mut [i8],
76) -> Result<()> {
77    let input_batches = input_dims.n as usize;
78    let input_h = input_dims.h as usize;
79    let input_w = input_dims.w as usize;
80    let channels = input_dims.c as usize;
81
82    let kernel_h = filter_dims.h as usize;
83    let kernel_w = filter_dims.w as usize;
84
85    let output_h = output_dims.h as usize;
86    let output_w = output_dims.w as usize;
87
88    for b in 0..input_batches {
89        for out_y in 0..output_h {
90            let base_y = out_y as i32 * pool_params.stride.h - pool_params.padding.h;
91            for out_x in 0..output_w {
92                let base_x = out_x as i32 * pool_params.stride.w - pool_params.padding.w;
93
94                for c in 0..channels {
95                    let mut sum: i32 = 0;
96                    let mut count: i32 = 0;
97
98                    for ky in 0..kernel_h {
99                        let in_y = base_y + ky as i32;
100                        if in_y >= 0 && in_y < input_dims.h {
101                            for kx in 0..kernel_w {
102                                let in_x = base_x + kx as i32;
103                                if in_x >= 0 && in_x < input_dims.w {
104                                    let in_idx = ((b * input_h + in_y as usize) * input_w
105                                        + in_x as usize)
106                                        * channels
107                                        + c;
108                                    sum += input[in_idx] as i32;
109                                    count += 1;
110                                }
111                            }
112                        }
113                    }
114
115                    let avg = if count > 0 {
116                        if sum >= 0 {
117                            (sum + count / 2) / count
118                        } else {
119                            (sum - count / 2) / count
120                        }
121                    } else {
122                        0
123                    };
124
125                    let final_val =
126                        clamp(avg, pool_params.activation.min, pool_params.activation.max);
127                    let out_idx = ((b * output_h + out_y) * output_w + out_x) * channels + c;
128                    output[out_idx] = final_val as i8;
129                }
130            }
131        }
132    }
133    Ok(())
134}
135
136/// Performs Max Pooling 2D for int16 tensors.
137pub fn max_pool_s16(
138    pool_params: &PoolParams,
139    filter_dims: &Tile,
140    input_dims: &Dims,
141    input: &[i16],
142    output_dims: &Dims,
143    output: &mut [i16],
144) -> Result<()> {
145    let input_batches = input_dims.n as usize;
146    let input_h = input_dims.h as usize;
147    let input_w = input_dims.w as usize;
148    let channels = input_dims.c as usize;
149
150    let kernel_h = filter_dims.h as usize;
151    let kernel_w = filter_dims.w as usize;
152
153    let output_h = output_dims.h as usize;
154    let output_w = output_dims.w as usize;
155
156    for b in 0..input_batches {
157        for out_y in 0..output_h {
158            let base_y = out_y as i32 * pool_params.stride.h - pool_params.padding.h;
159            for out_x in 0..output_w {
160                let base_x = out_x as i32 * pool_params.stride.w - pool_params.padding.w;
161
162                for c in 0..channels {
163                    let mut max_val = i16::MIN as i32;
164
165                    for ky in 0..kernel_h {
166                        let in_y = base_y + ky as i32;
167                        if in_y >= 0 && in_y < input_dims.h {
168                            for kx in 0..kernel_w {
169                                let in_x = base_x + kx as i32;
170                                if in_x >= 0 && in_x < input_dims.w {
171                                    let in_idx = ((b * input_h + in_y as usize) * input_w
172                                        + in_x as usize)
173                                        * channels
174                                        + c;
175                                    let val = input[in_idx] as i32;
176                                    if val > max_val {
177                                        max_val = val;
178                                    }
179                                }
180                            }
181                        }
182                    }
183
184                    max_val = clamp(
185                        max_val,
186                        pool_params.activation.min,
187                        pool_params.activation.max,
188                    );
189                    let out_idx = ((b * output_h + out_y) * output_w + out_x) * channels + c;
190                    output[out_idx] = max_val as i16;
191                }
192            }
193        }
194    }
195    Ok(())
196}
197
198/// Performs Average Pooling 2D for int16 tensors.
199pub fn avg_pool_s16(
200    pool_params: &PoolParams,
201    filter_dims: &Tile,
202    input_dims: &Dims,
203    input: &[i16],
204    output_dims: &Dims,
205    output: &mut [i16],
206) -> Result<()> {
207    let input_batches = input_dims.n as usize;
208    let input_h = input_dims.h as usize;
209    let input_w = input_dims.w as usize;
210    let channels = input_dims.c as usize;
211
212    let kernel_h = filter_dims.h as usize;
213    let kernel_w = filter_dims.w as usize;
214
215    let output_h = output_dims.h as usize;
216    let output_w = output_dims.w as usize;
217
218    for b in 0..input_batches {
219        for out_y in 0..output_h {
220            let base_y = out_y as i32 * pool_params.stride.h - pool_params.padding.h;
221            for out_x in 0..output_w {
222                let base_x = out_x as i32 * pool_params.stride.w - pool_params.padding.w;
223
224                for c in 0..channels {
225                    let mut sum: i64 = 0;
226                    let mut count: i64 = 0;
227
228                    for ky in 0..kernel_h {
229                        let in_y = base_y + ky as i32;
230                        if in_y >= 0 && in_y < input_dims.h {
231                            for kx in 0..kernel_w {
232                                let in_x = base_x + kx as i32;
233                                if in_x >= 0 && in_x < input_dims.w {
234                                    let in_idx = ((b * input_h + in_y as usize) * input_w
235                                        + in_x as usize)
236                                        * channels
237                                        + c;
238                                    sum += input[in_idx] as i64;
239                                    count += 1;
240                                }
241                            }
242                        }
243                    }
244
245                    let avg = if count > 0 {
246                        if sum >= 0 {
247                            (sum + count / 2) / count
248                        } else {
249                            (sum - count / 2) / count
250                        }
251                    } else {
252                        0
253                    };
254
255                    let final_val = clamp(
256                        avg as i32,
257                        pool_params.activation.min,
258                        pool_params.activation.max,
259                    );
260                    let out_idx = ((b * output_h + out_y) * output_w + out_x) * channels + c;
261                    output[out_idx] = final_val as i16;
262                }
263            }
264        }
265    }
266    Ok(())
267}
268
269#[cfg(test)]
270mod tests {
271    use super::*;
272    use crate::types::Activation;
273
274    #[test]
275    fn test_max_pool_s8() {
276        let pool_params = PoolParams {
277            stride: Tile::new(2, 2),
278            padding: Tile::new(0, 0),
279            activation: Activation::int8_unconstrained(),
280        };
281
282        let filter_dims = Tile::new(2, 2);
283        let input_dims = Dims::new(1, 4, 4, 1);
284        let input = [
285            1i8, 2i8, 5i8, 6i8, 3i8, 4i8, 7i8, 8i8, 9i8, 10i8, 13i8, 14i8, 11i8, 12i8, 15i8, 16i8,
286        ];
287
288        let output_dims = Dims::new(1, 2, 2, 1);
289        let mut output = [0i8; 4];
290
291        max_pool_s8(
292            &pool_params,
293            &filter_dims,
294            &input_dims,
295            &input,
296            &output_dims,
297            &mut output,
298        )
299        .unwrap();
300
301        assert_eq!(output, [4, 8, 12, 16]);
302    }
303}