| 179 | } |
| 180 | |
| 181 | const 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(); |