| 26 | |
| 27 | template <typename T0, typename T1 = T0> |
| 28 | void bn_forward_exec( |
| 29 | _megdnn_tensor_in src, _megdnn_tensor_in bn_scale, _megdnn_tensor_in bn_bias, |
| 30 | _megdnn_tensor_inout mean, _megdnn_tensor_inout variance, |
| 31 | _megdnn_tensor_out batch_mean, _megdnn_tensor_out batch_inv_variance, |
| 32 | _megdnn_tensor_out dst, param::BN param) { |
| 33 | size_t src_shape[4], dim_offset[4], param_pos = 0, src_pos = 0, batch_size = 1; |
| 34 | |
| 35 | float sigma_p, tmp, epsilon = (float)param.epsilon, denominator = 1.f; |
| 36 | |
| 37 | T0 *src_p = src.ptr<T0>(), *dst_p = dst.ptr<T0>(); |
| 38 | T1 *bn_scale_p = bn_scale.ptr<T1>(), *bn_bias_p = bn_bias.ptr<T1>(), |
| 39 | *mean_p = mean.ptr<T1>(), *variance_p = variance.ptr<T1>(), |
| 40 | *batch_mean_p = batch_mean.ptr<T1>(), |
| 41 | *batch_inv_variance_p = batch_inv_variance.ptr<T1>(); |
| 42 | |
| 43 | rep(i, src.layout.ndim) { |
| 44 | src_shape[i] = src.layout.shape[i]; |
| 45 | if (bn_scale.layout.shape[i] == 1) { |
| 46 | dim_offset[i] = 0; |
| 47 | batch_size *= src_shape[i]; |
| 48 | } else { |
| 49 | dim_offset[i] = 1; |
| 50 | } |
| 51 | } |
| 52 | |
| 53 | int curr_stride = 0; |
| 54 | for (int i = 3; i >= 0; --i) { |
| 55 | if (dim_offset[i] != 0) { |
| 56 | if (curr_stride == 0) { |
| 57 | dim_offset[i] = 1; |
| 58 | curr_stride = src_shape[i]; |
| 59 | } else { |
| 60 | dim_offset[i] = curr_stride; |
| 61 | curr_stride *= src_shape[i]; |
| 62 | } |
| 63 | } |
| 64 | } |
| 65 | denominator = 1.0 / batch_size; |
| 66 | |
| 67 | if (param.fwd_mode == param::BN::FwdMode::TRAINING) { |
| 68 | // Calculate the means of this batch (Mu) |
| 69 | memset(batch_mean_p, 0, batch_mean.layout.total_nr_elems() * sizeof(float)); |
| 70 | rep_4d(src_shape, dim_offset) batch_mean_p[param_pos] += src_p[src_pos]; |
| 71 | rep_4d_end |
| 72 | |
| 73 | rep(i, batch_mean.layout.total_nr_elems()) { |
| 74 | batch_mean_p[i] *= denominator; |
| 75 | if (!mean.layout.is_empty()) { |
| 76 | mean_p[i] = (1 - param.avg_factor) * mean_p[i] + |
| 77 | param.avg_factor * batch_mean_p[i]; |
| 78 | } |
| 79 | } |
| 80 | |
| 81 | // Calculate the variances of this batch (Sigma) |
| 82 | memset(batch_inv_variance_p, 0, |
| 83 | batch_inv_variance.layout.total_nr_elems() * sizeof(float)); |
| 84 | rep_4d(src_shape, dim_offset) sigma_p = |
| 85 | src_p[src_pos] - batch_mean_p[param_pos]; |
nothing calls this directly
no test coverage detected