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

Method save_model

models/mixmatch/mixmatch.py:270–288  ·  view source on GitHub ↗
(self, save_name, save_path)

Source from the content-addressed store, hash-verified

268 return {'eval/loss': total_loss / total_num, 'eval/top-1-acc': top1, 'eval/top-5-acc': top5}
269
270 def save_model(self, save_name, save_path):
271 if self.it < 1000000:
272 return
273 save_filename = os.path.join(save_path, save_name)
274 # copy EMA parameters to ema_model for saving with model as temp
275 self.model.eval()
276 self.ema.apply_shadow()
277 ema_model = deepcopy(self.model)
278 self.ema.restore()
279 self.model.train()
280
281 torch.save({'model': self.model.state_dict(),
282 'optimizer': self.optimizer.state_dict(),
283 'scheduler': self.scheduler.state_dict(),
284 'it': self.it + 1,
285 'ema_model': ema_model.state_dict()},
286 save_filename)
287
288 self.print_fn(f"model saved: {save_filename}")
289
290 def load_model(self, load_path):
291 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