use embedded_nn::{
convolution::{
convolve_1_x_n_s8, convolve_per_channel_s8, convolve_s8, depthwise_conv_per_channel_s8,
transpose_conv_s8,
},
float_ops::convolve_f32,
Activation, ConvParams, Dims, DwConvParams, Error, PerChannelQuantParams, PerTensorQuantParams,
Tile,
};
#[test]
fn test_convolve_s8_standard_per_tensor_variations() {
let conv_params = ConvParams {
input_offset: 128,
output_offset: -128,
stride: Tile::new(1, 1),
padding: Tile::new(0, 0),
dilation: Tile::new(1, 1),
activation: Activation::int8_unconstrained(),
};
let quant_params = PerTensorQuantParams::new(1073741824, 0);
let input_dims = Dims::new(1, 3, 3, 1);
let input = [10i8; 9];
let filter_dims = Dims::new(1, 3, 3, 1); let kernel = [1i8; 9];
let bias = [100i32];
let output_dims = Dims::new(1, 1, 1, 1);
let mut output = [0i8; 1];
convolve_s8(
&conv_params,
&quant_params,
&input_dims,
&input,
&filter_dims,
&kernel,
Some(&bias),
&output_dims,
&mut output,
)
.unwrap();
assert_eq!(output[0], 127);
}
#[test]
fn test_convolve_s8_stride_padding_dilation() {
let conv_params = ConvParams {
input_offset: 0,
output_offset: 0,
stride: Tile::new(2, 2),
padding: Tile::new(1, 1),
dilation: Tile::new(1, 1),
activation: Activation::new(0, 100),
};
let quant_params = PerTensorQuantParams::new(2147483647, 0);
let input_dims = Dims::new(1, 5, 5, 1);
let input = [1i8; 25];
let filter_dims = Dims::new(1, 3, 3, 1);
let kernel = [1i8; 9];
let output_dims = Dims::new(1, 3, 3, 1);
let mut output = [0i8; 9];
convolve_s8(
&conv_params,
&quant_params,
&input_dims,
&input,
&filter_dims,
&kernel,
None,
&output_dims,
&mut output,
)
.unwrap();
assert_eq!(output[0], 4);
assert_eq!(output[4], 9);
}
#[test]
fn test_convolve_s8_multi_batch_multi_channel() {
let conv_params = ConvParams {
input_offset: 0,
output_offset: 0,
stride: Tile::new(1, 1),
padding: Tile::new(0, 0),
dilation: Tile::new(1, 1),
activation: Activation::int8_unconstrained(),
};
let quant_params = PerTensorQuantParams::new(2147483647, 0);
let input_dims = Dims::new(2, 2, 2, 2); let input = [
1i8, 2i8, 3i8, 4i8, 5i8, 6i8, 7i8, 8i8, 9i8, 10i8, 11i8, 12i8, 13i8, 14i8, 15i8, 16i8, ];
let filter_dims = Dims::new(2, 2, 2, 2); let kernel = [1i8; 16];
let output_dims = Dims::new(2, 1, 1, 2);
let mut output = [0i8; 4];
convolve_s8(
&conv_params,
&quant_params,
&input_dims,
&input,
&filter_dims,
&kernel,
None,
&output_dims,
&mut output,
)
.unwrap();
assert_eq!(output[0], 36);
assert_eq!(output[1], 36);
assert_eq!(output[2], 100);
assert_eq!(output[3], 100);
}
#[test]
fn test_convolve_per_channel_s8_comprehensive() {
let conv_params = ConvParams {
input_offset: 0,
output_offset: 0,
stride: Tile::new(1, 1),
padding: Tile::new(0, 0),
dilation: Tile::new(1, 1),
activation: Activation::int8_unconstrained(),
};
let mults = [1073741824, 2147483647]; let shifts = [0, 0];
let quant_params = PerChannelQuantParams::new(&mults, &shifts);
let input_dims = Dims::new(1, 2, 2, 1);
let input = [10i8, 10i8, 10i8, 10i8];
let filter_dims = Dims::new(2, 2, 2, 1);
let kernel = [1i8; 8];
let output_dims = Dims::new(1, 1, 1, 2);
let mut output = [0i8; 2];
convolve_per_channel_s8(
&conv_params,
&quant_params,
&input_dims,
&input,
&filter_dims,
&kernel,
None,
&output_dims,
&mut output,
)
.unwrap();
assert_eq!(output[0], 20);
assert_eq!(output[1], 40);
}
#[test]
fn test_depthwise_conv_per_channel_s8_execution() {
let dw_params = DwConvParams {
input_offset: 0,
output_offset: 0,
ch_mult: 1,
stride: Tile::new(1, 1),
padding: Tile::new(0, 0),
dilation: Tile::new(1, 1),
activation: Activation::int8_unconstrained(),
};
let mults = [1073741824, 2147483647];
let shifts = [0, 0];
let quant_params = PerChannelQuantParams::new(&mults, &shifts);
let input_dims = Dims::new(1, 2, 2, 2);
let input = [10i8; 8];
let filter_dims = Dims::new(2, 2, 2, 1);
let kernel = [1i8; 8];
let output_dims = Dims::new(1, 1, 1, 2);
let mut output = [0i8; 2];
depthwise_conv_per_channel_s8(
&dw_params,
&quant_params,
&input_dims,
&input,
&filter_dims,
&kernel,
None,
&output_dims,
&mut output,
)
.unwrap();
assert_eq!(output[0], 20);
assert_eq!(output[1], 40);
}
#[test]
fn test_convolve_1_x_n_s8_temporal() {
let conv_params = ConvParams {
input_offset: 0,
output_offset: 0,
stride: Tile::new(1, 1),
padding: Tile::new(0, 0),
dilation: Tile::new(1, 1),
activation: Activation::int8_unconstrained(),
};
let quant_params = PerTensorQuantParams::new(2147483647, 0);
let input_dims = Dims::new(1, 1, 4, 1); let input = [1i8, 2i8, 3i8, 4i8];
let filter_dims = Dims::new(1, 1, 2, 1); let kernel = [1i8, 1i8];
let output_dims = Dims::new(1, 1, 3, 1);
let mut output = [0i8; 3];
convolve_1_x_n_s8(
&conv_params,
&quant_params,
&input_dims,
&input,
&filter_dims,
&kernel,
None,
&output_dims,
&mut output,
)
.unwrap();
assert_eq!(output, [3, 5, 7]);
}
#[test]
fn test_transpose_conv_s8_execution() {
let conv_params = ConvParams {
input_offset: 0,
output_offset: 0,
stride: Tile::new(2, 2),
padding: Tile::new(0, 0),
dilation: Tile::new(1, 1),
activation: Activation::int8_unconstrained(),
};
let mults = [2147483647];
let shifts = [0];
let quant_params = PerChannelQuantParams::new(&mults, &shifts);
let input_dims = Dims::new(1, 2, 2, 1);
let input = [1i8, 2i8, 3i8, 4i8];
let filter_dims = Dims::new(1, 2, 2, 1);
let kernel = [1i8, 1i8, 1i8, 1i8];
let output_dims = Dims::new(1, 4, 4, 1);
let mut output = [0i8; 16];
transpose_conv_s8(
&conv_params,
&quant_params,
&input_dims,
&input,
&filter_dims,
&kernel,
None,
&output_dims,
&mut output,
)
.unwrap();
assert_eq!(output[0], 1);
assert_eq!(output[5], 1);
assert_eq!(output[6], 2);
assert_eq!(output[10], 4);
}
#[test]
fn test_convolve_f32_execution() {
let input_dims = Dims::new(1, 3, 3, 1);
let input = [1.0f32; 9];
let filter_dims = Dims::new(1, 3, 3, 1);
let kernel = [2.0f32; 9];
let bias = [5.0f32];
let output_dims = Dims::new(1, 1, 1, 1);
let mut output = [0.0f32; 1];
convolve_f32(
Tile::new(1, 1),
Tile::new(0, 0),
Tile::new(1, 1),
&input_dims,
&input,
&filter_dims,
&kernel,
Some(&bias),
&output_dims,
&mut output,
)
.unwrap();
assert_eq!(output[0], 23.0);
}
#[test]
fn test_convolve_s8_error_paths() {
let conv_params = ConvParams {
input_offset: 0,
output_offset: 0,
stride: Tile::new(1, 1),
padding: Tile::new(0, 0),
dilation: Tile::new(1, 1),
activation: Activation::int8_unconstrained(),
};
let quant_params = PerTensorQuantParams::new(2147483647, 0);
let input_dims = Dims::new(1, 2, 2, 0); let input = [];
let filter_dims = Dims::new(1, 1, 1, 0);
let kernel = [];
let output_dims = Dims::new(1, 2, 2, 1);
let mut output = [0i8; 4];
let err = convolve_s8(
&conv_params,
&quant_params,
&input_dims,
&input,
&filter_dims,
&kernel,
None,
&output_dims,
&mut output,
);
assert_eq!(err, Err(Error::ArgumentError));
}