MCPcopy Create free account
hub / github.com/modelscope/modelscope / retrieve_modules_from_names

Function retrieve_modules_from_names

modelscope/utils/checkpoint.py:487–514  ·  view source on GitHub ↗
(model,
                                    names,
                                    prefix=None,
                                    add_prefix=False,
                                    remove_prefix=False)

Source from the content-addressed store, hash-verified

485 return missing_keys, unexpected_keys, mismatched_keys, error_msgs
486
487 def retrieve_modules_from_names(model,
488 names,
489 prefix=None,
490 add_prefix=False,
491 remove_prefix=False):
492 module_keys = set(['.'.join(key.split('.')[:-1]) for key in names])
493
494 # torch.nn.ParameterList is a special case where two parameter keywords
495 # are appended to the module name, *e.g.* bert.special_embeddings.0
496 module_keys = module_keys.union(
497 set([
498 '.'.join(key.split('.')[:-2]) for key in names
499 if key[-1].isdigit()
500 ]))
501
502 retrieved_modules = []
503 # retrieve all modules that has at least one missing weight name
504 for name, module in model.named_modules():
505 if remove_prefix:
506 name = '.'.join(
507 name.split('.')[1:]) if name.startswith(prefix) else name
508 elif add_prefix:
509 name = '.'.join([prefix, name]) if len(name) > 0 else prefix
510
511 if name in module_keys:
512 retrieved_modules.append(module)
513
514 return retrieved_modules
515
516 def _tie_or_clone_weights(output_embeddings,
517 input_embeddings,

Callers 1

_load_checkpointFunction · 0.85

Calls 1

appendMethod · 0.45

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…