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

Class Bn_Controller

train_utils.py:388–409  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

386
387
388class Bn_Controller:
389 def __init__(self):
390 """
391 freeze_bn and unfreeze_bn must appear in pairs
392 """
393 self.backup = {}
394
395 def freeze_bn(self, model):
396 assert self.backup == {}
397 for name, m in model.named_modules():
398 if isinstance(m, nn.SyncBatchNorm) or isinstance(m, nn.BatchNorm2d):
399 self.backup[name + '.running_mean'] = m.running_mean.data.clone()
400 self.backup[name + '.running_var'] = m.running_var.data.clone()
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 13

__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90
__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected