| 121 | } |
| 122 | |
| 123 | const std::vector<Tensor> CpuBatchNormForwardTraining( |
| 124 | const BatchNormHandle& bnh, const Tensor& x, const Tensor& bnScale, |
| 125 | const Tensor& bnBias, Tensor& running_mean, Tensor& running_var) { |
| 126 | CHECK_EQ(x.device()->lang(), kCpp); |
| 127 | Tensor y; |
| 128 | y.ResetLike(x); |
| 129 | |
| 130 | // mean and var for local batch |
| 131 | Tensor mean; |
| 132 | mean.ResetLike(running_mean); |
| 133 | Tensor var; |
| 134 | var.ResetLike(running_var); |
| 135 | |
| 136 | // combine scale and bias to construct weight tensor in required format for |
| 137 | // backward |
| 138 | Tensor w = get_bn_weight_from(bnScale, bnBias); |
| 139 | |
| 140 | y.device()->Exec( |
| 141 | [y, mean, var, w, x, &running_mean, &running_var, |
| 142 | &bnh](Context* ctx) mutable { |
| 143 | auto eng = ctx->dnnl_engine; |
| 144 | using namespace dnnl; |
| 145 | |
| 146 | auto x_mem = memory(bnh.x_md, eng, x.block()->mutable_data()); |
| 147 | auto y_mem = memory(bnh.x_md, eng, y.block()->mutable_data()); |
| 148 | auto m_mem = memory(bnh.bn_fwd_training_pd->mean_desc(), eng, |
| 149 | mean.block()->mutable_data()); |
| 150 | auto v_mem = memory(bnh.bn_fwd_training_pd->variance_desc(), eng, |
| 151 | var.block()->mutable_data()); |
| 152 | auto w_mem = memory(bnh.bn_fwd_training_pd->weights_desc(), eng, |
| 153 | w.block()->mutable_data()); |
| 154 | |
| 155 | batch_normalization_forward(*bnh.bn_fwd_training_pd) |
| 156 | .execute(ctx->dnnl_stream, {{DNNL_ARG_SRC, x_mem}, |
| 157 | {DNNL_ARG_DST, y_mem}, |
| 158 | {DNNL_ARG_SCALE_SHIFT, w_mem}, |
| 159 | {DNNL_ARG_MEAN, m_mem}, |
| 160 | {DNNL_ARG_VARIANCE, v_mem}}); |
| 161 | ctx->dnnl_stream.wait(); |
| 162 | |
| 163 | // local implemented running mean as mkldnn does not support it yet: |
| 164 | // https://github.com/intel/mkl-dnn/issues/371 |
| 165 | // https://github.com/intel/mkl-dnn/issues/517 |
| 166 | // https://arxiv.org/pdf/1502.03167.pdf |
| 167 | auto s = x.shape(); |
| 168 | s[1] = 1; |
| 169 | float p = Product(s); // for unbiased variance |
| 170 | running_mean = running_mean * (1 - bnh.factor) + mean * bnh.factor; |
| 171 | running_var = |
| 172 | running_var * (1 - bnh.factor) + var * (p / (p - 1)) * bnh.factor; |
| 173 | }, |
| 174 | {x.block(), w.block(), running_mean.block(), running_var.block()}, |
| 175 | {y.block(), running_mean.block(), running_var.block(), mean.block(), |
| 176 | var.block()}, "CpuBatchNormForwardTraining"); |
| 177 | |
| 178 | return {y, mean, var}; |
| 179 | } |
| 180 | |