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

Method save_model

models/pseudolabel/pseudolabel.py:244–262  ·  view source on GitHub ↗
(self, save_name, save_path)

Source from the content-addressed store, hash-verified

242 'eval/precision': precision, 'eval/recall': recall, 'eval/F1': F1, 'eval/AUC': AUC}
243
244 def save_model(self, save_name, save_path):
245 if self.it < 1000000:
246 return
247 save_filename = os.path.join(save_path, save_name)
248 # copy EMA parameters to ema_model for saving with model as temp
249 self.model.eval()
250 self.ema.apply_shadow()
251 ema_model = self.model.state_dict()
252 self.ema.restore()
253 self.model.train()
254
255 torch.save({'model': self.model.state_dict(),
256 'optimizer': self.optimizer.state_dict(),
257 'scheduler': self.scheduler.state_dict(),
258 'it': self.it + 1,
259 'ema_model': ema_model},
260 save_filename)
261
262 self.print_fn(f"model saved: {save_filename}")
263
264 def load_model(self, load_path):
265 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