1use crate::support::{clamp, requantize};
4use crate::types::{
5 ConvParams, Dims, DwConvParams, Error, PerChannelQuantParams, PerTensorQuantParams, Result,
6};
7
8pub fn convolve_s8(
10 conv_params: &ConvParams,
11 quant_params: &PerTensorQuantParams,
12 input_dims: &Dims,
13 input: &[i8],
14 filter_dims: &Dims,
15 kernel: &[i8],
16 bias: Option<&[i32]>,
17 output_dims: &Dims,
18 output: &mut [i8],
19) -> Result<()> {
20 let input_batches = input_dims.n as usize;
21 let input_h = input_dims.h as usize;
22 let input_w = input_dims.w as usize;
23 let input_c = input_dims.c as usize;
24
25 let kernel_h = filter_dims.h as usize;
26 let kernel_w = filter_dims.w as usize;
27 let kernel_c = filter_dims.c as usize;
28
29 let output_h = output_dims.h as usize;
30 let output_w = output_dims.w as usize;
31 let output_c = output_dims.c as usize;
32
33 if input_c == 0 || output_c == 0 {
34 return Err(Error::ArgumentError);
35 }
36
37 let groups = input_c / kernel_c;
38 let output_c_per_group = output_c / groups;
39
40 for b in 0..input_batches {
41 for out_y in 0..output_h {
42 let base_y = out_y as i32 * conv_params.stride.h - conv_params.padding.h;
43 for out_x in 0..output_w {
44 let base_x = out_x as i32 * conv_params.stride.w - conv_params.padding.w;
45
46 for g in 0..groups {
47 for out_ch_idx in 0..output_c_per_group {
48 let out_c = g * output_c_per_group + out_ch_idx;
49 let mut acc: i32 = match bias {
50 Some(b_slice) => b_slice[out_c],
51 None => 0,
52 };
53
54 for ky in 0..kernel_h {
55 let in_y = base_y + ky as i32 * conv_params.dilation.h;
56 if in_y >= 0 && in_y < input_dims.h {
57 for kx in 0..kernel_w {
58 let in_x = base_x + kx as i32 * conv_params.dilation.w;
59 if in_x >= 0 && in_x < input_dims.w {
60 let in_idx_base = ((b * input_h + in_y as usize) * input_w
61 + in_x as usize)
62 * input_c
63 + g * kernel_c;
64 let ker_idx_base =
65 ((out_c * kernel_h + ky) * kernel_w + kx) * kernel_c;
66
67 for ic in 0..kernel_c {
68 let lhs = input[in_idx_base + ic] as i32
69 + conv_params.input_offset;
70 let rhs = kernel[ker_idx_base + ic] as i32;
71 acc += lhs * rhs;
72 }
73 }
74 }
75 }
76 }
77
78 acc = requantize(acc, quant_params.multiplier, quant_params.shift);
79 acc += conv_params.output_offset;
80 acc = clamp(acc, conv_params.activation.min, conv_params.activation.max);
81
82 let out_idx =
83 ((b * output_h + out_y) * output_w + out_x) * output_c + out_c;
84 output[out_idx] = acc as i8;
85 }
86 }
87 }
88 }
89 }
90
91 Ok(())
92}
93
94pub fn convolve_per_channel_s8(
96 conv_params: &ConvParams,
97 quant_params: &PerChannelQuantParams,
98 input_dims: &Dims,
99 input: &[i8],
100 filter_dims: &Dims,
101 kernel: &[i8],
102 bias: Option<&[i32]>,
103 output_dims: &Dims,
104 output: &mut [i8],
105) -> Result<()> {
106 let input_batches = input_dims.n as usize;
107 let input_h = input_dims.h as usize;
108 let input_w = input_dims.w as usize;
109 let input_c = input_dims.c as usize;
110
111 let kernel_h = filter_dims.h as usize;
112 let kernel_w = filter_dims.w as usize;
113 let kernel_c = filter_dims.c as usize;
114
115 let output_h = output_dims.h as usize;
116 let output_w = output_dims.w as usize;
117 let output_c = output_dims.c as usize;
118
119 if input_c == 0 || output_c == 0 {
120 return Err(Error::ArgumentError);
121 }
122
123 let groups = input_c / kernel_c;
124 let output_c_per_group = output_c / groups;
125
126 for b in 0..input_batches {
127 for out_y in 0..output_h {
128 let base_y = out_y as i32 * conv_params.stride.h - conv_params.padding.h;
129 for out_x in 0..output_w {
130 let base_x = out_x as i32 * conv_params.stride.w - conv_params.padding.w;
131
132 for g in 0..groups {
133 for out_ch_idx in 0..output_c_per_group {
134 let out_c = g * output_c_per_group + out_ch_idx;
135 let mut acc: i32 = match bias {
136 Some(b_slice) => b_slice[out_c],
137 None => 0,
138 };
139
140 for ky in 0..kernel_h {
141 let in_y = base_y + ky as i32 * conv_params.dilation.h;
142 if in_y >= 0 && in_y < input_dims.h {
143 for kx in 0..kernel_w {
144 let in_x = base_x + kx as i32 * conv_params.dilation.w;
145 if in_x >= 0 && in_x < input_dims.w {
146 let in_idx_base = ((b * input_h + in_y as usize) * input_w
147 + in_x as usize)
148 * input_c
149 + g * kernel_c;
150 let ker_idx_base =
151 ((out_c * kernel_h + ky) * kernel_w + kx) * kernel_c;
152
153 for ic in 0..kernel_c {
154 let lhs = input[in_idx_base + ic] as i32
155 + conv_params.input_offset;
156 let rhs = kernel[ker_idx_base + ic] as i32;
157 acc += lhs * rhs;
158 }
159 }
160 }
161 }
162 }
163
164 let mult = quant_params.multiplier[out_c];
165 let shift = quant_params.shift[out_c];
166
167 acc = requantize(acc, mult, shift);
168 acc += conv_params.output_offset;
169 acc = clamp(acc, conv_params.activation.min, conv_params.activation.max);
170
171 let out_idx =
172 ((b * output_h + out_y) * output_w + out_x) * output_c + out_c;
173 output[out_idx] = acc as i8;
174 }
175 }
176 }
177 }
178 }
179
180 Ok(())
181}
182
183pub fn depthwise_conv_per_channel_s8(
185 dw_params: &DwConvParams,
186 quant_params: &PerChannelQuantParams,
187 input_dims: &Dims,
188 input: &[i8],
189 filter_dims: &Dims,
190 kernel: &[i8],
191 bias: Option<&[i32]>,
192 output_dims: &Dims,
193 output: &mut [i8],
194) -> Result<()> {
195 let input_batches = input_dims.n as usize;
196 let input_h = input_dims.h as usize;
197 let input_w = input_dims.w as usize;
198 let input_c = input_dims.c as usize;
199
200 let kernel_h = filter_dims.h as usize;
201 let kernel_w = filter_dims.w as usize;
202
203 let output_h = output_dims.h as usize;
204 let output_w = output_dims.w as usize;
205 let output_c = output_dims.c as usize;
206
207 let ch_mult = dw_params.ch_mult as usize;
208
209 for b in 0..input_batches {
210 for out_y in 0..output_h {
211 let base_y = out_y as i32 * dw_params.stride.h - dw_params.padding.h;
212 for out_x in 0..output_w {
213 let base_x = out_x as i32 * dw_params.stride.w - dw_params.padding.w;
214
215 for out_c in 0..output_c {
216 let in_c = out_c / ch_mult;
217
218 let mut acc: i32 = match bias {
219 Some(b_slice) => b_slice[out_c],
220 None => 0,
221 };
222
223 for ky in 0..kernel_h {
224 let in_y = base_y + ky as i32 * dw_params.dilation.h;
225 if in_y >= 0 && in_y < input_dims.h {
226 for kx in 0..kernel_w {
227 let in_x = base_x + kx as i32 * dw_params.dilation.w;
228 if in_x >= 0 && in_x < input_dims.w {
229 let in_idx = ((b * input_h + in_y as usize) * input_w
230 + in_x as usize)
231 * input_c
232 + in_c;
233 let ker_idx = ((ky * kernel_w + kx) * output_c) + out_c;
234
235 let lhs = input[in_idx] as i32 + dw_params.input_offset;
236 let rhs = kernel[ker_idx] as i32;
237 acc += lhs * rhs;
238 }
239 }
240 }
241 }
242
243 let mult = quant_params.multiplier[out_c];
244 let shift = quant_params.shift[out_c];
245
246 acc = requantize(acc, mult, shift);
247 acc += dw_params.output_offset;
248 acc = clamp(acc, dw_params.activation.min, dw_params.activation.max);
249
250 let out_idx = ((b * output_h + out_y) * output_w + out_x) * output_c + out_c;
251 output[out_idx] = acc as i8;
252 }
253 }
254 }
255 }
256
257 Ok(())
258}
259
260pub fn transpose_conv_s8(
262 conv_params: &ConvParams,
263 quant_params: &PerChannelQuantParams,
264 input_dims: &Dims,
265 input: &[i8],
266 filter_dims: &Dims,
267 kernel: &[i8],
268 bias: Option<&[i32]>,
269 output_dims: &Dims,
270 output: &mut [i8],
271) -> Result<()> {
272 let input_batches = input_dims.n as usize;
273 let input_h = input_dims.h as usize;
274 let input_w = input_dims.w as usize;
275 let input_c = input_dims.c as usize;
276
277 let kernel_h = filter_dims.h as usize;
278 let kernel_w = filter_dims.w as usize;
279
280 let output_h = output_dims.h as usize;
281 let output_w = output_dims.w as usize;
282 let output_c = output_dims.c as usize;
283
284 let stride_y = conv_params.stride.h as usize;
285 let stride_x = conv_params.stride.w as usize;
286 let pad_y = conv_params.padding.h as usize;
287 let pad_x = conv_params.padding.w as usize;
288
289 let mut accum_buffer = [0i32; 1024]; let scratch_size = output_h * output_w * output_c;
291
292 for b in 0..input_batches {
293 for out_idx in 0..scratch_size {
294 let out_c = out_idx % output_c;
295 let b_val = match bias {
296 Some(b_slice) => b_slice[out_c],
297 None => 0,
298 };
299 if out_idx < accum_buffer.len() {
300 accum_buffer[out_idx] = b_val;
301 }
302 }
303
304 for in_y in 0..input_h {
306 for in_x in 0..input_w {
307 for ky in 0..kernel_h {
308 let out_y = in_y * stride_y + ky;
309 if out_y >= pad_y && out_y < output_h + pad_y {
310 let actual_out_y = out_y - pad_y;
311 for kx in 0..kernel_w {
312 let out_x = in_x * stride_x + kx;
313 if out_x >= pad_x && out_x < output_w + pad_x {
314 let actual_out_x = out_x - pad_x;
315
316 for out_c in 0..output_c {
317 for in_c in 0..input_c {
318 let in_idx = ((b * input_h + in_y) * input_w + in_x)
319 * input_c
320 + in_c;
321 let ker_idx = ((out_c * kernel_h + ky) * kernel_w + kx)
322 * input_c
323 + in_c;
324
325 let lhs = input[in_idx] as i32 + conv_params.input_offset;
326 let rhs = kernel[ker_idx] as i32;
327
328 let buf_idx = (actual_out_y * output_w + actual_out_x)
329 * output_c
330 + out_c;
331 if buf_idx < accum_buffer.len() {
332 accum_buffer[buf_idx] += lhs * rhs;
333 }
334 }
335 }
336 }
337 }
338 }
339 }
340 }
341 }
342
343 for out_y in 0..output_h {
345 for out_x in 0..output_w {
346 for out_c in 0..output_c {
347 let buf_idx = (out_y * output_w + out_x) * output_c + out_c;
348 let acc = if buf_idx < accum_buffer.len() {
349 accum_buffer[buf_idx]
350 } else {
351 match bias {
352 Some(b_slice) => b_slice[out_c],
353 None => 0,
354 }
355 };
356
357 let mult = quant_params.multiplier[out_c];
358 let shift = quant_params.shift[out_c];
359
360 let req = requantize(acc, mult, shift);
361 let final_val = clamp(
362 req + conv_params.output_offset,
363 conv_params.activation.min,
364 conv_params.activation.max,
365 );
366
367 let out_idx = ((b * output_h + out_y) * output_w + out_x) * output_c + out_c;
368 if out_idx < output.len() {
369 output[out_idx] = final_val as i8;
370 }
371 }
372 }
373 }
374 }
375
376 Ok(())
377}
378
379pub fn convolve_1_x_n_s8(
381 conv_params: &ConvParams,
382 quant_params: &PerTensorQuantParams,
383 input_dims: &Dims,
384 input: &[i8],
385 filter_dims: &Dims,
386 kernel: &[i8],
387 bias: Option<&[i32]>,
388 output_dims: &Dims,
389 output: &mut [i8],
390) -> Result<()> {
391 convolve_s8(
393 conv_params,
394 quant_params,
395 input_dims,
396 input,
397 filter_dims,
398 kernel,
399 bias,
400 output_dims,
401 output,
402 )
403}
404
405#[cfg(test)]
406mod tests {
407 use super::*;
408 use crate::types::{Activation, Tile};
409
410 #[test]
411 fn test_convolve_s8_simple() {
412 let conv_params = ConvParams {
413 input_offset: 0,
414 output_offset: 0,
415 stride: Tile::new(1, 1),
416 padding: Tile::new(0, 0),
417 dilation: Tile::new(1, 1),
418 activation: Activation::int8_unconstrained(),
419 };
420
421 let quant_params = PerTensorQuantParams::new(1073741824, 0); let input_dims = Dims::new(1, 3, 3, 1);
424 let input = [1i8, 2i8, 3i8, 4i8, 5i8, 6i8, 7i8, 8i8, 9i8];
425
426 let filter_dims = Dims::new(1, 2, 2, 1); let kernel = [1i8, 0i8, 0i8, 1i8];
428
429 let output_dims = Dims::new(1, 2, 2, 1);
430 let mut output = [0i8; 4];
431
432 convolve_s8(
433 &conv_params,
434 &quant_params,
435 &input_dims,
436 &input,
437 &filter_dims,
438 &kernel,
439 None,
440 &output_dims,
441 &mut output,
442 )
443 .unwrap();
444
445 assert_eq!(output, [3, 4, 6, 7]);
446 }
447}