(model,
names,
prefix=None,
add_prefix=False,
remove_prefix=False)
| 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, |
no test coverage detected
searching dependent graphs…