class DecoderWrapper(torch.nn.Module): def __init__(self, decoder_model): super().__init__() self.decoder = decoder_model def forward(self, input_ids, encoder_hidden_states): return self.decoder( input_ids=input_ids, encoder_hidden_states=encoder_hidden_states, use_cache=False, output_attentions=False, output_hidden_states=False, return_dict=False ) def load_model(path=EXPORT_PATH, mode=None): model = get_model() encoder = model.encoder decoder = model.decoder return encoder, DecoderWrapper(decoder)