| 1112 | # PyTorch's `_load_from_state_dict` does not copy parameters in a module's descendants |
| 1113 | # so we need to apply the function recursively. |
| 1114 | def load(module: nn.Module, prefix=""): |
| 1115 | local_metadata = {} if metadata is None else metadata.get(prefix[:-1], {}) |
| 1116 | module._load_from_state_dict( |
| 1117 | state_dict, |
| 1118 | prefix, |
| 1119 | local_metadata, |
| 1120 | True, |
| 1121 | missing_keys, |
| 1122 | unexpected_keys, |
| 1123 | error_msgs, |
| 1124 | ) |
| 1125 | for name, child in module._modules.items(): |
| 1126 | if child is not None: |
| 1127 | load(child, prefix + name + ".") |
| 1128 | |
| 1129 | # Make sure we are able to load base models as well as derived models (with heads) |
| 1130 | start_prefix = "" |