MCPcopy Create free account
hub / github.com/TorchSSL/TorchSSL / unfreeze_bn

Method unfreeze_bn

train_utils.py:403–409  ·  view source on GitHub ↗
(self, model)

Source from the content-addressed store, hash-verified

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 = {}

Callers 6

trainMethod · 0.80
trainMethod · 0.80
trainMethod · 0.80
trainMethod · 0.80
trainMethod · 0.80
trainMethod · 0.80

Calls

no outgoing calls

Tested by

no test coverage detected