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

Function bn_backward_exec

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

Source from the content-addressed store, hash-verified

115
116template <typename T0, typename T1 = T0>
117void bn_backward_exec(
118 _megdnn_tensor_in x_in, _megdnn_tensor_in dy_in,
119 _megdnn_tensor_in saved_batch_mean, _megdnn_tensor_in saved_batch_inv_variance,
120 _megdnn_tensor_in bn_scale, _megdnn_tensor_out d_bn_scale,
121 _megdnn_tensor_out d_bn_bias, _megdnn_tensor_out dx_out,
122 const WorkspaceBundle& bundle) {
123 size_t src_shape[4], dim_offset[4], param_pos = 0, src_pos = 0, batch_size = 1,
124 param_size = bn_scale.layout.total_nr_elems();
125
126 float denominator = 1.f;
127
128 T0 *x = x_in.ptr<T0>(), *dx = dx_out.ptr<T0>(), *dy = dy_in.ptr<T0>();
129 T1 *gamma = bn_scale.ptr<T1>(), *mu = saved_batch_mean.ptr<T1>(),
130 *ivar = saved_batch_inv_variance.ptr<T1>(), *dgamma = d_bn_scale.ptr<T1>(),
131 *dbeta = d_bn_bias.ptr<T1>();
132
133 rep(i, dy_in.layout.ndim) {
134 src_shape[i] = dy_in.layout.shape[i];
135 if (bn_scale.layout.shape[i] == 1) {
136 dim_offset[i] = 0;
137 batch_size *= src_shape[i];
138 } else {
139 dim_offset[i] = 1;
140 }
141 }
142
143 int curr_stride = 0;
144 for (int i = 3; i >= 0; --i) {
145 if (dim_offset[i] != 0) {
146 if (curr_stride == 0) {
147 dim_offset[i] = 1;
148 curr_stride = src_shape[i];
149 } else {
150 dim_offset[i] = curr_stride;
151 curr_stride *= src_shape[i];
152 }
153 }
154 }
155 denominator = 1.0 / batch_size;
156
157 // step1. dbeta, dgamma
158 memset(dbeta, 0, param_size * sizeof(T1));
159 memset(dgamma, 0, param_size * sizeof(T1));
160 rep_4d(src_shape, dim_offset) float xhat =
161 (x[src_pos] - mu[param_pos]) * ivar[param_pos];
162 dbeta[param_pos] += dy[src_pos];
163 dgamma[param_pos] += dy[src_pos] * xhat;
164 rep_4d_end
165
166 // step2. dxhat = dy * gamma
167 float* dxhat = static_cast<float*>(bundle.get(0));
168 rep_4d(src_shape, dim_offset) dxhat[src_pos] = dy[src_pos] * gamma[param_pos];
169 rep_4d_end
170
171 // step3. dvar = sigma[dxhat * xmu] * [-1/2 * (ivar)^(3/2)]
172 // dmu = sigma[dxhat] * (-ivar)
173 float* dvar = static_cast<float*>(bundle.get(1));
174 float* dmu = static_cast<float*>(bundle.get(2));

Callers

nothing calls this directly

Calls 3

repFunction · 0.70
total_nr_elemsMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected