(module)
| 36 | |
| 37 | |
| 38 | def _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 | |
| 59 | def _demo_mm_inputs(input_shape, num_classes): |
no outgoing calls
no test coverage detected