MCPcopy Create free account
hub / github.com/QWTforGithub/T2LDM / adapt_model_from_string

Function adapt_model_from_string

timm/models/helpers.py:279–326  ·  view source on GitHub ↗
(parent_module, model_string)

Source from the content-addressed store, hash-verified

277
278
279def 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
329def adapt_model_from_file(parent_module, model_variant):

Callers 1

adapt_model_from_fileFunction · 0.85

Calls 3

extract_layerFunction · 0.85
set_layerFunction · 0.85
LinearClass · 0.85

Tested by

no test coverage detected