def capture_model(model, path=EXPORT_PATH): # weights only weights_path = os.path.join(EXPORT_PATH, "weights.pth") torch.save(model.state_dict(), weights_path) def load_model(path=EXPORT_PATH, mode=None): if mode == 'weights': model = get_model() weights_path = os.path.join(path,"weights.pth") state_dict = torch.load(weights_path, map_location="cpu") model.load_state_dict(state_dict) return model.encoder, DecoderWrapper(model.decoder) else: model = get_model() return model.encoder, DecoderWrapper(model.decoder)