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

Function CpuBatchNormForwardTraining

src/model/operation/batchnorm.cc:123–179  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

121}
122
123const 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

Callers 1

TESTFunction · 0.85

Calls 9

get_bn_weight_fromFunction · 0.85
ProductFunction · 0.85
langMethod · 0.80
deviceMethod · 0.80
ExecMethod · 0.80
mutable_dataMethod · 0.80
shapeMethod · 0.80
blockMethod · 0.45
waitMethod · 0.45

Tested by 1

TESTFunction · 0.68