Save networks. Args: net (nn.Module | list[nn.Module]): Network(s) to be saved. net_label (str): Network label. current_iter (int): Current iter number. param_key (str | list[str]): The parameter key(s) to save network. Default
(self, net, net_label, current_iter, param_key='params')
| 210 | |
| 211 | @master_only |
| 212 | def save_network(self, net, net_label, current_iter, param_key='params'): |
| 213 | """Save networks. |
| 214 | |
| 215 | Args: |
| 216 | net (nn.Module | list[nn.Module]): Network(s) to be saved. |
| 217 | net_label (str): Network label. |
| 218 | current_iter (int): Current iter number. |
| 219 | param_key (str | list[str]): The parameter key(s) to save network. |
| 220 | Default: 'params'. |
| 221 | """ |
| 222 | if current_iter == -1: |
| 223 | current_iter = 'latest' |
| 224 | save_filename = f'{net_label}_{current_iter}.pth' |
| 225 | save_path = os.path.join(self.opt['path']['models'], save_filename) |
| 226 | |
| 227 | net = net if isinstance(net, list) else [net] |
| 228 | param_key = param_key if isinstance(param_key, list) else [param_key] |
| 229 | assert len(net) == len(param_key), 'The lengths of net and param_key should be the same.' |
| 230 | |
| 231 | save_dict = {} |
| 232 | for net_, param_key_ in zip(net, param_key): |
| 233 | net_ = self.get_bare_model(net_) |
| 234 | state_dict = net_.state_dict() |
| 235 | for key, param in state_dict.items(): |
| 236 | if key.startswith('module.'): # remove unnecessary 'module.' |
| 237 | key = key[7:] |
| 238 | state_dict[key] = param.cpu() |
| 239 | save_dict[param_key_] = state_dict |
| 240 | |
| 241 | # avoid occasional writing errors |
| 242 | retry = 3 |
| 243 | while retry > 0: |
| 244 | try: |
| 245 | torch.save(save_dict, save_path) |
| 246 | except Exception as e: |
| 247 | logger = get_root_logger() |
| 248 | logger.warning(f'Save model error: {e}, remaining retry times: {retry - 1}') |
| 249 | time.sleep(1) |
| 250 | else: |
| 251 | break |
| 252 | finally: |
| 253 | retry -= 1 |
| 254 | if retry == 0: |
| 255 | logger.warning(f'Still cannot save {save_path}. Just ignore it.') |
| 256 | # raise IOError(f'Cannot save {save_path}.') |
| 257 | |
| 258 | def _print_different_keys_loading(self, crt_net, load_net, strict=True): |
| 259 | """Print keys with different name or different size when loading models. |
no test coverage detected