MCPcopy Create free account
hub / github.com/apache/singa / CpuBatchNormBackwardx

Function CpuBatchNormBackwardx

src/model/operation/batchnorm.cc:181–256  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

179}
180
181const std::vector<Tensor> CpuBatchNormBackwardx(
182 const BatchNormHandle& bnh, const Tensor& y, const Tensor& dy,
183 const Tensor& x, const Tensor& bnScale, const Tensor& bnBias,
184 const Tensor& mean, const Tensor& var) {
185 CHECK_EQ(x.device()->lang(), kCpp);
186 CHECK_EQ(y.device()->lang(), kCpp);
187 CHECK_EQ(dy.device()->lang(), kCpp);
188 CHECK_EQ(mean.device()->lang(), kCpp);
189 CHECK_EQ(var.device()->lang(), kCpp);
190 CHECK_EQ(bnScale.device()->lang(), kCpp);
191 CHECK_EQ(bnBias.device()->lang(), kCpp);
192
193 Tensor dx;
194 dx.ResetLike(dy);
195
196 // combine scale and bias to construct weight tensor in required format for
197 // backward
198 Tensor w = get_bn_weight_from(bnScale, bnBias);
199
200 // Tensor dw(Shape{bnScale.Size(), 2});
201 Tensor dw;
202 dw.ResetLike(w);
203
204 dx.device()->Exec(
205 [w, dw, dx, dy, x, y, mean, var, &bnh](Context* ctx) mutable {
206 auto eng = ctx->dnnl_engine;
207 using namespace dnnl;
208
209 auto x_mem = memory(bnh.x_md, eng, x.block()->mutable_data());
210 auto dx_mem = memory(bnh.x_md, eng, dx.block()->mutable_data());
211 auto y_mem = memory(bnh.x_md, eng, y.block()->mutable_data());
212 auto dy_mem = memory(bnh.x_md, eng, dy.block()->mutable_data());
213
214 auto m_mem = memory(bnh.bn_fwd_training_pd->mean_desc(), eng,
215 mean.block()->mutable_data());
216 auto v_mem = memory(bnh.bn_fwd_training_pd->variance_desc(), eng,
217 var.block()->mutable_data());
218 auto w_mem = memory(bnh.bn_fwd_training_pd->weights_desc(), eng,
219 w.block()->mutable_data());
220
221 auto bn_bwd_d = batch_normalization_backward::desc(
222 prop_kind::backward, bnh.x_md, bnh.x_md, bnh.epsilon,
223 normalization_flags::use_scale_shift);
224 auto bn_bwd_pd = batch_normalization_backward::primitive_desc(
225 bn_bwd_d, eng, *bnh.bn_fwd_training_pd);
226
227 auto dw_mem = memory(bn_bwd_pd.diff_weights_desc(), eng,
228 dw.block()->mutable_data());
229
230 batch_normalization_backward(bn_bwd_pd).execute(
231 ctx->dnnl_stream, {{DNNL_ARG_SRC, x_mem},
232 {DNNL_ARG_DIFF_SRC, dx_mem},
233 {DNNL_ARG_DIFF_DST, dy_mem},
234 {DNNL_ARG_MEAN, m_mem},
235 {DNNL_ARG_VARIANCE, v_mem},
236 {DNNL_ARG_DIFF_SCALE_SHIFT, dw_mem},
237 {DNNL_ARG_SCALE_SHIFT, w_mem}});
238 ctx->dnnl_stream.wait();

Callers 1

TESTFunction · 0.85

Calls 11

get_bn_weight_fromFunction · 0.85
CopyDataToFromFunction · 0.85
langMethod · 0.80
deviceMethod · 0.80
ExecMethod · 0.80
mutable_dataMethod · 0.80
shapeMethod · 0.80
nDimMethod · 0.80
blockMethod · 0.45
waitMethod · 0.45
SizeMethod · 0.45

Tested by 1

TESTFunction · 0.68