1use crate::support::clamp;
4use crate::types::{Dims, PoolParams, Result, Tile};
5
6pub 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
68pub 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
136pub 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
198pub 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}