| 156 | } |
| 157 | |
| 158 | const 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 | |