MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / bn_forward_exec

Function bn_forward_exec

dnn/src/naive/batch_normalization/opr_impl.cpp:28–114  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

26
27template <typename T0, typename T1 = T0>
28void 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];

Callers

nothing calls this directly

Calls 4

repFunction · 0.70
sqrtFunction · 0.50
total_nr_elemsMethod · 0.45
is_emptyMethod · 0.45

Tested by

no test coverage detected