def capture_model(model, path=EXPORT_PATH): encoder = model.encoder decoder = DecoderWrapper(model.decoder) # define dynamic dimensions batch = torch.export.Dim("batch") seq_len = torch.export.Dim("seq_len", min=2, max=MAX_SEQ_LEN) # export encoder # sample input example = torch.randn(4, 3, 224, 224) encoder_export = torch.export.export( encoder, (example,), dynamic_shapes=((batch, torch.export.Dim.STATIC, torch.export.Dim.STATIC, torch.export.Dim.STATIC), ) ) torch.export.save(encoder_export, os.path.join(path, "encoder.pt2")) # export decoder # get sample input for decoder encoder_hidden_states = encoder_export.module()(example)[0] decoder_input_ids = torch.ones((4, MAX_SEQ_LEN), dtype=torch.long)*START_ID decoder_export = torch.export.export( decoder, (decoder_input_ids, encoder_hidden_states), dynamic_shapes={ 'input_ids': (batch,seq_len), 'encoder_hidden_states': (batch, torch.export.Dim.STATIC, torch.export.Dim.STATIC) } ) torch.export.save(decoder_export, os.path.join(path, "decoder.pt2")) 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 == 'export': encoder_path = os.path.join(path, "encoder.pt2") decoder_path = os.path.join(path, "decoder.pt2") encoder = load(encoder_path).module() decoder = load(decoder_path).module() return encoder, decoder else: model = get_model() return model.encoder, DecoderWrapper(model.decoder)