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

Method Backward

src/model/layer/batchnorm.cc:158–245  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

156}
157
158const std::pair<Tensor, vector<Tensor>> BatchNorm::Backward(
159 int flag, const Tensor& grad) {
160 Tensor dy = grad.Clone();
161 dy.Reshape(Shape{grad.shape(0), grad.Size() / grad.shape(0)});
162 Tensor xnorm = buf_.top();
163 buf_.pop();
164 Tensor var = buf_.top();
165 buf_.pop();
166 Tensor mean = buf_.top();
167 buf_.pop();
168 Tensor input = buf_.top();
169 buf_.pop();
170
171 Tensor dx;
172 vector<Tensor> param_grad;
173
174 if ((flag & kTrain) == kTrain) {
175 if (is_2d_) {
176 // gxnrom
177 Tensor gxnorm = dy.Clone();
178 MultRow(bnScale_, &gxnorm);
179 // gvar
180 Tensor tmp = var.Clone();
181 tmp += 1e-6f;
182 tmp = Pow(var, -1.5f);
183 tmp *= -0.5f;
184
185 Tensor tmpx = input.Clone();
186 SubRow(mean, &tmpx);
187
188 tmpx = tmpx * gxnorm;
189 MultRow(tmp, &tmpx);
190 Tensor gvar;
191 gvar.ResetLike(var);
192 SumRows(tmpx, &gvar);
193 // gmean
194 tmp = var.Clone();
195 tmp += 1e-6f;
196 tmp = Pow(tmp, -0.5f);
197 tmp *= -1.0f;
198 Tensor tmpx_r;
199 tmpx_r.ResetLike(tmp);
200 SumRows(gxnorm, &tmpx_r);
201 Tensor gmean = tmpx_r * tmp;
202
203 tmpx = input.Clone();
204 SubRow(mean, &tmpx);
205 SumRows(tmpx, &tmp);
206 tmp *= -2.0f / input.shape(0);
207 tmp = tmp * gvar;
208 gmean = gmean + tmp;
209 // dx
210 tmp = var.Clone();
211 tmp += 1e-6f;
212 tmp = Pow(tmp, -0.5f);
213 dx = gxnorm.Clone();
214 MultRow(tmp, &dx);
215

Callers

nothing calls this directly

Calls 8

MultRowFunction · 0.85
SubRowFunction · 0.85
SumRowsFunction · 0.85
AddRowFunction · 0.85
shapeMethod · 0.80
PowFunction · 0.50
CloneMethod · 0.45
SizeMethod · 0.45

Tested by

no test coverage detected