(parent_module, model_string)
| 277 | |
| 278 | |
| 279 | def adapt_model_from_string(parent_module, model_string): |
| 280 | separator = '***' |
| 281 | state_dict = {} |
| 282 | lst_shape = model_string.split(separator) |
| 283 | for k in lst_shape: |
| 284 | k = k.split(':') |
| 285 | key = k[0] |
| 286 | shape = k[1][1:-1].split(',') |
| 287 | if shape[0] != '': |
| 288 | state_dict[key] = [int(i) for i in shape] |
| 289 | |
| 290 | new_module = deepcopy(parent_module) |
| 291 | for n, m in parent_module.named_modules(): |
| 292 | old_module = extract_layer(parent_module, n) |
| 293 | if isinstance(old_module, nn.Conv2d) or isinstance(old_module, Conv2dSame): |
| 294 | if isinstance(old_module, Conv2dSame): |
| 295 | conv = Conv2dSame |
| 296 | else: |
| 297 | conv = nn.Conv2d |
| 298 | s = state_dict[n + '.weight'] |
| 299 | in_channels = s[1] |
| 300 | out_channels = s[0] |
| 301 | g = 1 |
| 302 | if old_module.groups > 1: |
| 303 | in_channels = out_channels |
| 304 | g = in_channels |
| 305 | new_conv = conv( |
| 306 | in_channels=in_channels, out_channels=out_channels, kernel_size=old_module.kernel_size, |
| 307 | bias=old_module.bias is not None, padding=old_module.padding, dilation=old_module.dilation, |
| 308 | groups=g, stride=old_module.stride) |
| 309 | set_layer(new_module, n, new_conv) |
| 310 | if isinstance(old_module, nn.BatchNorm2d): |
| 311 | new_bn = nn.BatchNorm2d( |
| 312 | num_features=state_dict[n + '.weight'][0], eps=old_module.eps, momentum=old_module.momentum, |
| 313 | affine=old_module.affine, track_running_stats=True) |
| 314 | set_layer(new_module, n, new_bn) |
| 315 | if isinstance(old_module, nn.Linear): |
| 316 | # FIXME extra checks to ensure this is actually the FC classifier layer and not a diff Linear layer? |
| 317 | num_features = state_dict[n + '.weight'][1] |
| 318 | new_fc = Linear( |
| 319 | in_features=num_features, out_features=old_module.out_features, bias=old_module.bias is not None) |
| 320 | set_layer(new_module, n, new_fc) |
| 321 | if hasattr(new_module, 'num_features'): |
| 322 | new_module.num_features = num_features |
| 323 | new_module.eval() |
| 324 | parent_module.eval() |
| 325 | |
| 326 | return new_module |
| 327 | |
| 328 | |
| 329 | def adapt_model_from_file(parent_module, model_variant): |
no test coverage detected