(self, model)
| 401 | self.backup[name + '.num_batches_tracked'] = m.num_batches_tracked.data.clone() |
| 402 | |
| 403 | def unfreeze_bn(self, model): |
| 404 | for name, m in model.named_modules(): |
| 405 | if isinstance(m, nn.SyncBatchNorm) or isinstance(m, nn.BatchNorm2d): |
| 406 | m.running_mean.data = self.backup[name + '.running_mean'] |
| 407 | m.running_var.data = self.backup[name + '.running_var'] |
| 408 | m.num_batches_tracked.data = self.backup[name + '.num_batches_tracked'] |
| 409 | self.backup = {} |