| 29 | // A slow but straightforward implementation of batch normalization. |
| 30 | template <typename T1, typename T2> |
| 31 | void ReferenceBatchNorm(const Tensor& input, const float input_min, |
| 32 | const float input_max, const Tensor& mean, |
| 33 | float mean_min, float mean_max, const Tensor& var, |
| 34 | float var_min, float var_max, const Tensor& beta, |
| 35 | float beta_min, float beta_max, const Tensor& gamma, |
| 36 | float gamma_min, float gamma_max, |
| 37 | float variance_epsilon, bool scale_after_normalization, |
| 38 | Tensor* output, float* output_min, float* output_max) { |
| 39 | auto input_flat = input.flat<T1>(); |
| 40 | auto mean_flat = mean.flat<T1>(); |
| 41 | auto var_flat = var.flat<T1>(); |
| 42 | auto beta_flat = beta.flat<T1>(); |
| 43 | auto gamma_flat = gamma.flat<T1>(); |
| 44 | auto output_flat = output->flat<T2>(); |
| 45 | |
| 46 | const int depth = mean.dim_size(0); |
| 47 | const int row_count = input_flat.size() / depth; |
| 48 | |
| 49 | *output_min = std::numeric_limits<float>::max(); |
| 50 | *output_max = std::numeric_limits<float>::lowest(); |
| 51 | for (int pass = 0; pass < 2; ++pass) { |
| 52 | const bool is_range_pass = (pass == 0); |
| 53 | for (int row_index = 0; row_index < row_count; ++row_index) { |
| 54 | for (int channel = 0; channel < depth; ++channel) { |
| 55 | const int input_index = (row_index * depth) + channel; |
| 56 | const float input_value = |
| 57 | QuantizedToFloat(input_flat(input_index), input_min, input_max); |
| 58 | const float mean_value = |
| 59 | QuantizedToFloat(mean_flat(channel), mean_min, mean_max); |
| 60 | const float var_value = |
| 61 | QuantizedToFloat(var_flat(channel), var_min, var_max); |
| 62 | const float beta_value = |
| 63 | QuantizedToFloat(beta_flat(channel), beta_min, beta_max); |
| 64 | const float gamma_value = |
| 65 | QuantizedToFloat(gamma_flat(channel), gamma_min, gamma_max); |
| 66 | float output_value; |
| 67 | if (scale_after_normalization) { |
| 68 | output_value = (((input_value - mean_value) / |
| 69 | sqrtf(var_value + variance_epsilon)) * |
| 70 | gamma_value) + |
| 71 | beta_value; |
| 72 | } else { |
| 73 | output_value = ((input_value - mean_value) / |
| 74 | sqrtf(var_value + variance_epsilon)) + |
| 75 | beta_value; |
| 76 | } |
| 77 | if (is_range_pass) { |
| 78 | *output_min = std::min(output_value, *output_min); |
| 79 | *output_max = std::max(output_value, *output_max); |
| 80 | } else { |
| 81 | output_flat(input_index) = |
| 82 | FloatToQuantized<T2>(output_value, *output_min, *output_max); |
| 83 | } |
| 84 | } |
| 85 | } |
| 86 | } |
| 87 | } |
| 88 |
nothing calls this directly
no test coverage detected