Designed to work with network_to_half. BatchNorm layers need parameters in single precision. Find all layers and convert them back to float. This can't be done with built in .apply as that function will apply fn to all modules, parameters, and buffers. Thus we wouldn't be ab
(module)
| 204 | |
| 205 | |
| 206 | def BN_convert_float(module): |
| 207 | ''' |
| 208 | Designed to work with network_to_half. |
| 209 | BatchNorm layers need parameters in single precision. |
| 210 | Find all layers and convert them back to float. This can't |
| 211 | be done with built in .apply as that function will apply |
| 212 | fn to all modules, parameters, and buffers. Thus we wouldn't |
| 213 | be able to guard the float conversion based on the module type. |
| 214 | ''' |
| 215 | if isinstance(module, torch.nn.modules.batchnorm._BatchNorm): |
| 216 | module.float() |
| 217 | for child in module.children(): |
| 218 | BN_convert_float(child) |
| 219 | return module |
| 220 | |
| 221 | |
| 222 | def network_to_half(network): |