| 126 | self.load_network(self.model, load_path, strict=False) |
| 127 | |
| 128 | def _register_hooks(self): |
| 129 | module_name_list = [] |
| 130 | hook_list = [] |
| 131 | name_list = [] |
| 132 | def get_module_from_name(name): |
| 133 | name_parts = name.split('.')[:-1] |
| 134 | module_name = 'self.model' |
| 135 | |
| 136 | for part in name_parts: |
| 137 | if part == 'mask_token' or part == 'weight' or part == 'bias': |
| 138 | continue |
| 139 | if part.isdigit(): |
| 140 | module_name += f'[{part}]' |
| 141 | else: |
| 142 | module_name += f'.{part}' |
| 143 | return module_name, '.'.join(name_parts) |
| 144 | |
| 145 | for name,param in self.model.named_parameters(): |
| 146 | module_name,name = get_module_from_name(name) |
| 147 | |
| 148 | # print(module_name) |
| 149 | if module_name != 'self.model': |
| 150 | if len(module_name_list)==0 or module_name_list[-1] != module_name: |
| 151 | module_name_list.append(module_name) |
| 152 | name_list.append(name) |
| 153 | module = eval(module_name) |
| 154 | hook_list.append(Hook_back_loop(module, name)) |
| 155 | |
| 156 | return hook_list |
| 157 | |
| 158 | def build_test_loaders(self): |
| 159 | test_loaders = [] |