MCPcopy Create free account
hub / github.com/TorchSSL/TorchSSL / save_model

Method save_model

models/pimodel/pimodel.py:227–245  ·  view source on GitHub ↗
(self, save_name, save_path)

Source from the content-addressed store, hash-verified

225 return {'eval/loss': total_loss / total_num, 'eval/top-1-acc': top1, 'eval/top-5-acc': top5}
226
227 def save_model(self, save_name, save_path):
228 if self.it < 1000000:
229 return
230 save_filename = os.path.join(save_path, save_name)
231 # copy EMA parameters to ema_model for saving with model as temp
232 self.model.eval()
233 self.ema.apply_shadow()
234 ema_model = deepcopy(self.model)
235 self.ema.restore()
236 self.model.train()
237
238 torch.save({'model': self.model.state_dict(),
239 'optimizer': self.optimizer.state_dict(),
240 'scheduler': self.scheduler.state_dict(),
241 'it': self.it + 1,
242 'ema_model': ema_model.state_dict()},
243 save_filename)
244
245 self.print_fn(f"model saved: {save_filename}")
246
247 def load_model(self, load_path):
248 checkpoint = torch.load(load_path)

Callers 1

trainMethod · 0.95

Calls 3

apply_shadowMethod · 0.80
restoreMethod · 0.80
trainMethod · 0.45

Tested by

no test coverage detected