import torch NUM_TOKENS = 1024 MAX_SEQ_LEN = 256 PAD_ID = 0 START_ID = 1 END_ID = 2 # Set up an image-to-text model. def get_model(): # import transformers utilities from transformers import ( VisionEncoderDecoderModel, VisionEncoderDecoderConfig, AutoConfig ) config = VisionEncoderDecoderConfig.from_encoder_decoder_configs( encoder_config=AutoConfig.for_model("vit"), # vit encoder decoder_config=AutoConfig.for_model("gpt2") # gpt2 decoder ) config.decoder.vocab_size = NUM_TOKENS config.decoder.use_cache = False config.decoder_start_token_id = START_ID config.pad_token_id = PAD_ID config.eos_token_id = END_ID config.max_length = MAX_SEQ_LEN model = VisionEncoderDecoderModel(config=config) model.encoder.pooler = None # remove unused pooler model.eval() # prepare the model for evaluation return model