def generate_token(decoder, encoder_hidden_states, sequence): outputs = decoder( sequence, encoder_hidden_states, torch.arange(sequence.shape[1], device=sequence.device) ) logits = outputs[0][:, -1, :] return torch.argmax(logits, dim=-1, keepdim=True) class DecoderWrapper(torch.nn.Module): def __init__(self, decoder_model): super().__init__() self.decoder = decoder_model def forward(self, input_ids, encoder_hidden_states, cache_position): return self.decoder( input_ids=input_ids, cache_position=cache_position, encoder_hidden_states=encoder_hidden_states, use_cache=False, output_attentions=False, output_hidden_states=False, return_dict=False ) 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 tensor 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_args = ( decoder_input_ids, encoder_hidden_states, torch.arange(MAX_SEQ_LEN) ) dynamic_shapes = { 'input_ids': (batch,seq_len), 'encoder_hidden_states': (batch, torch.export.Dim.STATIC, torch.export.Dim.STATIC), 'cache_position': (seq_len,), } decoder_export = torch.export.export( decoder, decoder_args, dynamic_shapes=dynamic_shapes ) torch.export.save(decoder_export, os.path.join(path, "decoder.pt2"))