1use crate::support::{clamp, requantize};
4use crate::types::{Dims, FcParams, PerChannelQuantParams, PerTensorQuantParams, Result};
5
6pub fn fully_connected_s8(
8 fc_params: &FcParams,
9 quant_params: &PerTensorQuantParams,
10 input_dims: &Dims,
11 input: &[i8],
12 filter_dims: &Dims,
13 kernel: &[i8],
14 bias: Option<&[i32]>,
15 output_dims: &Dims,
16 output: &mut [i8],
17) -> Result<()> {
18 let batches = input_dims.n as usize;
19 let accum_depth = filter_dims.n as usize;
20 let output_depth = output_dims.c as usize;
21
22 for b in 0..batches {
23 let input_batch = &input[b * accum_depth..(b + 1) * accum_depth];
24 let output_batch = &mut output[b * output_depth..(b + 1) * output_depth];
25
26 for out_c in 0..output_depth {
27 let mut acc: i32 = match bias {
28 Some(b_slice) => b_slice[out_c],
29 None => 0,
30 };
31
32 let kernel_row = &kernel[out_c * accum_depth..(out_c + 1) * accum_depth];
33
34 for i in 0..accum_depth {
35 let lhs = input_batch[i] as i32 + fc_params.input_offset;
36 let rhs = kernel_row[i] as i32 + fc_params.filter_offset;
37 acc += lhs * rhs;
38 }
39
40 acc = requantize(acc, quant_params.multiplier, quant_params.shift);
41 acc += fc_params.output_offset;
42 acc = clamp(acc, fc_params.activation.min, fc_params.activation.max);
43
44 output_batch[out_c] = acc as i8;
45 }
46 }
47 Ok(())
48}
49
50pub fn fully_connected_per_channel_s8(
52 fc_params: &FcParams,
53 quant_params: &PerChannelQuantParams,
54 input_dims: &Dims,
55 input: &[i8],
56 filter_dims: &Dims,
57 kernel: &[i8],
58 bias: Option<&[i32]>,
59 output_dims: &Dims,
60 output: &mut [i8],
61) -> Result<()> {
62 let batches = input_dims.n as usize;
63 let accum_depth = filter_dims.n as usize;
64 let output_depth = output_dims.c as usize;
65
66 for b in 0..batches {
67 let input_batch = &input[b * accum_depth..(b + 1) * accum_depth];
68 let output_batch = &mut output[b * output_depth..(b + 1) * output_depth];
69
70 for out_c in 0..output_depth {
71 let mut acc: i32 = match bias {
72 Some(b_slice) => b_slice[out_c],
73 None => 0,
74 };
75
76 let kernel_row = &kernel[out_c * accum_depth..(out_c + 1) * accum_depth];
77
78 for i in 0..accum_depth {
79 let lhs = input_batch[i] as i32 + fc_params.input_offset;
80 let rhs = kernel_row[i] as i32 + fc_params.filter_offset;
81 acc += lhs * rhs;
82 }
83
84 let mult = quant_params.multiplier[out_c];
85 let shift = quant_params.shift[out_c];
86
87 acc = requantize(acc, mult, shift);
88 acc += fc_params.output_offset;
89 acc = clamp(acc, fc_params.activation.min, fc_params.activation.max);
90
91 output_batch[out_c] = acc as i8;
92 }
93 }
94 Ok(())
95}
96
97pub fn fully_connected_s16(
99 fc_params: &FcParams,
100 quant_params: &PerTensorQuantParams,
101 input_dims: &Dims,
102 input: &[i16],
103 filter_dims: &Dims,
104 kernel: &[i8],
105 bias: Option<&[i64]>,
106 output_dims: &Dims,
107 output: &mut [i16],
108) -> Result<()> {
109 let batches = input_dims.n as usize;
110 let accum_depth = filter_dims.n as usize;
111 let output_depth = output_dims.c as usize;
112
113 for b in 0..batches {
114 let input_batch = &input[b * accum_depth..(b + 1) * accum_depth];
115 let output_batch = &mut output[b * output_depth..(b + 1) * output_depth];
116
117 for out_c in 0..output_depth {
118 let mut acc: i64 = match bias {
119 Some(b_slice) => b_slice[out_c],
120 None => 0,
121 };
122
123 let kernel_row = &kernel[out_c * accum_depth..(out_c + 1) * accum_depth];
124
125 for i in 0..accum_depth {
126 let lhs = input_batch[i] as i64;
127 let rhs = kernel_row[i] as i64;
128 acc += lhs * rhs;
129 }
130
131 let req = requantize(
132 (acc >> 15) as i32,
133 quant_params.multiplier,
134 quant_params.shift,
135 );
136 let final_val = clamp(req, fc_params.activation.min, fc_params.activation.max);
137
138 output_batch[out_c] = final_val as i16;
139 }
140 }
141 Ok(())
142}
143
144pub fn batch_matmul_s8(
148 fc_params: &FcParams,
149 quant_params: &PerTensorQuantParams,
150 lhs_dims: &Dims,
151 input_lhs: &[i8],
152 rhs_dims: &Dims,
153 input_rhs: &[i8],
154 output_dims: &Dims,
155 output: &mut [i8],
156) -> Result<()> {
157 let batches = output_dims.n as usize;
158 let rows = lhs_dims.h as usize;
159 let cols = rhs_dims.c as usize;
160 let accum_dim = lhs_dims.w as usize;
161
162 for b in 0..batches {
163 let lhs_b_idx = if lhs_dims.n == 1 { 0 } else { b };
164 let rhs_b_idx = if rhs_dims.n == 1 { 0 } else { b };
165
166 for i in 0..rows {
167 for j in 0..cols {
168 let mut acc: i32 = 0;
169
170 for k in 0..accum_dim {
171 let lhs_idx = (lhs_b_idx * rows + i) * accum_dim + k;
172 let rhs_idx = (rhs_b_idx * accum_dim + k) * cols + j;
173
174 let lhs_val = input_lhs[lhs_idx] as i32 + fc_params.input_offset;
175 let rhs_val = input_rhs[rhs_idx] as i32 + fc_params.filter_offset;
176 acc += lhs_val * rhs_val;
177 }
178
179 let req = requantize(acc, quant_params.multiplier, quant_params.shift);
180 let final_val = clamp(
181 req + fc_params.output_offset,
182 fc_params.activation.min,
183 fc_params.activation.max,
184 );
185
186 let out_idx = (b * rows + i) * cols + j;
187 if out_idx < output.len() {
188 output[out_idx] = final_val as i8;
189 }
190 }
191 }
192 }
193
194 Ok(())
195}
196
197pub fn batch_matmul_s16(
199 fc_params: &FcParams,
200 quant_params: &PerTensorQuantParams,
201 lhs_dims: &Dims,
202 input_lhs: &[i16],
203 rhs_dims: &Dims,
204 input_rhs: &[i16],
205 output_dims: &Dims,
206 output: &mut [i16],
207) -> Result<()> {
208 let batches = output_dims.n as usize;
209 let rows = lhs_dims.h as usize;
210 let cols = rhs_dims.c as usize;
211 let accum_dim = lhs_dims.w as usize;
212
213 for b in 0..batches {
214 let lhs_b_idx = if lhs_dims.n == 1 { 0 } else { b };
215 let rhs_b_idx = if rhs_dims.n == 1 { 0 } else { b };
216
217 for i in 0..rows {
218 for j in 0..cols {
219 let mut acc: i64 = 0;
220
221 for k in 0..accum_dim {
222 let lhs_idx = (lhs_b_idx * rows + i) * accum_dim + k;
223 let rhs_idx = (rhs_b_idx * accum_dim + k) * cols + j;
224
225 let lhs_val = input_lhs[lhs_idx] as i64;
226 let rhs_val = input_rhs[rhs_idx] as i64;
227 acc += lhs_val * rhs_val;
228 }
229
230 let req = requantize(
231 (acc >> 15) as i32,
232 quant_params.multiplier,
233 quant_params.shift,
234 );
235 let final_val = clamp(req, fc_params.activation.min, fc_params.activation.max);
236
237 let out_idx = (b * rows + i) * cols + j;
238 if out_idx < output.len() {
239 output[out_idx] = final_val as i16;
240 }
241 }
242 }
243 }
244
245 Ok(())
246}
247
248#[cfg(test)]
249mod tests {
250 use super::*;
251 use crate::types::Activation;
252
253 #[test]
254 fn test_fully_connected_s8() {
255 let fc_params = FcParams {
256 input_offset: 0,
257 filter_offset: 0,
258 output_offset: 0,
259 activation: Activation::int8_unconstrained(),
260 };
261 let quant_params = PerTensorQuantParams::new(1073741824, 0); let input_dims = Dims::new(1, 1, 1, 3);
263 let input = [2i8, 4i8, 6i8];
264 let filter_dims = Dims::new(3, 1, 1, 2); let kernel = [
266 1i8, 2i8, 3i8, 4i8, 5i8, 6i8, ];
269 let bias = [0i32, 0i32];
270 let output_dims = Dims::new(1, 1, 1, 2);
271 let mut output = [0i8; 2];
272
273 fully_connected_s8(
274 &fc_params,
275 &quant_params,
276 &input_dims,
277 &input,
278 &filter_dims,
279 &kernel,
280 Some(&bias),
281 &output_dims,
282 &mut output,
283 )
284 .unwrap();
285
286 assert_eq!(output[0], 14);
287 assert_eq!(output[1], 32);
288 }
289
290 #[test]
291 fn test_batch_matmul_s8() {
292 let fc_params = FcParams {
293 input_offset: 0,
294 filter_offset: 0,
295 output_offset: 0,
296 activation: Activation::int8_unconstrained(),
297 };
298 let quant_params = PerTensorQuantParams::new(1073741824, 0); let lhs_dims = Dims::new(1, 2, 2, 0); let input_lhs = [1i8, 2i8, 3i8, 4i8];
302
303 let rhs_dims = Dims::new(1, 2, 0, 2); let input_rhs = [5i8, 6i8, 7i8, 8i8];
305
306 let output_dims = Dims::new(1, 2, 0, 2);
307 let mut output = [0i8; 4];
308
309 batch_matmul_s8(
310 &fc_params,
311 &quant_params,
312 &lhs_dims,
313 &input_lhs,
314 &rhs_dims,
315 &input_rhs,
316 &output_dims,
317 &mut output,
318 )
319 .unwrap();
320
321 assert_eq!(output, [10, 11, 22, 25]);
326 }
327}