MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / ReferenceBatchNorm

Function ReferenceBatchNorm

tensorflow/core/kernels/quantized_batch_norm_op.cc:31–87  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

29// A slow but straightforward implementation of batch normalization.
30template <typename T1, typename T2>
31void 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

Callers

nothing calls this directly

Calls 5

QuantizedToFloatFunction · 0.85
maxFunction · 0.50
minFunction · 0.50
dim_sizeMethod · 0.45
sizeMethod · 0.45

Tested by

no test coverage detected