MCPcopy Create free account
hub / github.com/DragonisCV/RAM / save_network

Method save_network

ram/models/base_model.py:212–256  ·  view source on GitHub ↗

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')

Source from the content-addressed store, hash-verified

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.

Callers 4

saveMethod · 0.80
saveMethod · 0.80
saveMethod · 0.80
saveMethod · 0.80

Calls 3

get_bare_modelMethod · 0.95
get_root_loggerFunction · 0.90
saveMethod · 0.45

Tested by

no test coverage detected