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) encoder = model.encoder decoder = DecoderWrapper(model.decoder) # torchscript encoder using trace example = torch.randn(1, 3, 224, 224) encoder_jit = torch.jit.trace(encoder, example) # optionally apply jit.freeze optimization encoder_jit = torch.jit.freeze(encoder_jit) encoder_path = os.path.join(path, "encoder.pt") torch.jit.save(encoder_jit, encoder_path) try: # torchscript decoder using scripting decoder_jit = torch.jit.script(decoder) # optionally apply jit.freeze optimization decoder_jit = torch.jit.freeze(decoder_jit) decoder_path = os.path.join(path, "decoder.pt") torch.jit.save(decoder_jit, decoder_path) except Exception as e: print(f'torch.jit.script(model.decoder) failed\n{e}') 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) elif mode == 'torchscript': encoder_path = os.path.join(path, "encoder.pt") decoder_path = os.path.join(path, "decoder.pt") encoder = torch.jit.load(encoder_path) decoder = torch.jit.load(decoder_path) # optionally apply target-device optimization encoder = torch.jit.optimize_for_inference(encoder) decoder = torch.jit.optimize_for_inference(decoder) return encoder, decoder else: model = get_model() return model.encoder, DecoderWrapper(model.decoder)