MCPcopy Create free account
hub / github.com/alinlab/SelfPatch / _convert_batchnorm

Function _convert_batchnorm

segmentation/tools/pytorch2torchscript.py:38–56  ·  view source on GitHub ↗
(module)

Source from the content-addressed store, hash-verified

36
37
38def _convert_batchnorm(module):
39 module_output = module
40 if isinstance(module, torch.nn.SyncBatchNorm):
41 module_output = torch.nn.BatchNorm2d(module.num_features, module.eps,
42 module.momentum, module.affine,
43 module.track_running_stats)
44 if module.affine:
45 module_output.weight.data = module.weight.data.clone().detach()
46 module_output.bias.data = module.bias.data.clone().detach()
47 # keep requires_grad unchanged
48 module_output.weight.requires_grad = module.weight.requires_grad
49 module_output.bias.requires_grad = module.bias.requires_grad
50 module_output.running_mean = module.running_mean
51 module_output.running_var = module.running_var
52 module_output.num_batches_tracked = module.num_batches_tracked
53 for name, child in module.named_children():
54 module_output.add_module(name, _convert_batchnorm(child))
55 del module
56 return module_output
57
58
59def _demo_mm_inputs(input_shape, num_classes):

Callers 1

Calls

no outgoing calls

Tested by

no test coverage detected