esp_nn_depthwise_conv_ansi.c 4.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100
  1. // Copyright 2020-2021 Espressif Systems (Shanghai) PTE LTD
  2. //
  3. // Licensed under the Apache License, Version 2.0 (the "License");
  4. // you may not use this file except in compliance with the License.
  5. // You may obtain a copy of the License at
  6. //
  7. // http://www.apache.org/licenses/LICENSE-2.0
  8. //
  9. // Unless required by applicable law or agreed to in writing, software
  10. // distributed under the License is distributed on an "AS IS" BASIS,
  11. // WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
  12. // See the License for the specific language governing permissions and
  13. // limitations under the License.
  14. #include <esp_nn_defs.h>
  15. #include <common_functions.h>
  16. int esp_nn_get_depthwise_conv_scratch_size_ansi(const data_dims_t *input_dims,
  17. const data_dims_t *filter_dims,
  18. const data_dims_t *output_dims,
  19. const dw_conv_params_t *conv_params)
  20. {
  21. return 0;
  22. }
  23. void esp_nn_set_depthwise_conv_scratch_buf_ansi(const void *buf)
  24. {
  25. }
  26. void esp_nn_depthwise_conv_s8_ansi(const data_dims_t *input_dims,
  27. const int8_t *input_data,
  28. const data_dims_t *filter_dims,
  29. const int8_t *filter_data,
  30. const int32_t *bias,
  31. const data_dims_t *output_dims,
  32. int8_t *out_data,
  33. const dw_conv_params_t *conv_params,
  34. const quant_data_t *quant_data)
  35. {
  36. const uint16_t input_wd = input_dims->width;
  37. const uint16_t input_ht = input_dims->height;
  38. const uint16_t channels = input_dims->channels;
  39. const int32_t input_offset = conv_params->in_offset;
  40. const int32_t out_offset = conv_params->out_offset;
  41. const uint16_t pad_wd = conv_params->padding.width;
  42. const uint16_t pad_ht = conv_params->padding.height;
  43. const uint16_t stride_wd = conv_params->stride.width;
  44. const uint16_t stride_ht = conv_params->stride.height;
  45. const uint16_t filter_wd = filter_dims->width;
  46. const uint16_t filter_ht = filter_dims->height;
  47. const uint16_t out_wd = output_dims->width;
  48. const uint16_t out_ht = output_dims->height;
  49. const int32_t *out_shift = quant_data->shift;
  50. const int32_t *out_mult = quant_data->mult;
  51. const int32_t activation_min = conv_params->activation.min;
  52. const int32_t activation_max = conv_params->activation.max;
  53. const uint16_t ch_mult = conv_params->ch_mult;
  54. int out_idx = 0;
  55. for (int out_y = 0; out_y < out_ht; out_y++) { //height loop
  56. const int16_t base_y = (out_y * stride_ht) - pad_ht;
  57. for (int out_x = 0; out_x < out_wd; out_x++) { //width_loop
  58. const int16_t base_x = (out_x * stride_wd) - pad_wd;
  59. for (int ch_idx = 0; ch_idx < channels; ch_idx++) {//channel_loop
  60. for (int ch_mult_idx = 0; ch_mult_idx < ch_mult; ch_mult_idx++) {
  61. int32_t result = 0;
  62. const int out_ch_idx = ch_mult_idx + ch_idx * ch_mult;
  63. /* Select filter so as the point doesn't lie outside block */
  64. int filter_y_start = max(0, -base_y);
  65. int filter_x_start = max(0, -base_x);
  66. int filter_y_end = min(filter_ht, input_ht - base_y);
  67. int filter_x_end = min(filter_wd, input_wd - base_x);
  68. for (int filter_y_idx = filter_y_start; filter_y_idx < filter_y_end; filter_y_idx++) {
  69. const int32_t idx_y = base_y + filter_y_idx;
  70. for (int filter_x_idx = filter_x_start; filter_x_idx < filter_x_end; filter_x_idx++) {
  71. const int32_t idx_x = base_x + filter_x_idx;
  72. int32_t input_index = (idx_y * input_wd + idx_x) * channels + ch_idx;
  73. int32_t filter_index = (filter_y_idx * filter_wd + filter_x_idx) * (channels * ch_mult) + out_ch_idx;
  74. int32_t input_val = input_data[input_index] + input_offset;
  75. int32_t filter_val = filter_data[filter_index];
  76. result += input_val * filter_val;
  77. }
  78. }
  79. if (bias) {
  80. result += bias[out_ch_idx];
  81. }
  82. result = esp_nn_multiply_by_quantized_mult(result, out_mult[out_ch_idx], out_shift[out_ch_idx]);
  83. result += out_offset;
  84. result = max(result, activation_min);
  85. result = min(result, activation_max);
  86. out_data[out_idx++] = result;
  87. }
  88. }
  89. }
  90. }
  91. }