Views
No views yet
1import onnxruntime as ort
2from transformers import AutoTokenizer,AutoImageProcessor
3from PIL import Image
4import numpy as np
5
6# load the ONNX models (encoder and decoder)
7encoder_onnx_path = 'models/rgb_language_cap_onnx/encoder_model.onnx' # load from local path
8decoder_onnx_path = 'models/rgb_language_cap_onnx/decoder_model.onnx' # load from local path
9encoder_session = ort.InferenceSession(encoder_onnx_path, providers=["CPUExecutionProvider"])
10decoder_session = ort.InferenceSession(decoder_onnx_path, providers=["CPUExecutionProvider"])
11
12# load the tokenizer and image processor
13model_id = "models/rgb_language_cap_onnx"
14processor = AutoImageProcessor.from_pretrained(model_id)
15tokenizer = AutoTokenizer.from_pretrained(model_id)
16
17# load image
18image_path = "img2.jpg"
19image = Image.open(image_path)
20inputs = processor(images=image, return_tensors="np").pixel_values
21
22# run encoder model
23encoder_outputs = encoder_session.run(
24 None,
25 {"pixel_values": inputs}
26)
27
28# extract the encoder hidden states (encoder outputs)
29encoder_hidden_states = encoder_outputs[0]
30
31# prepare decoder inputs
32decoder_input_ids = np.array([[tokenizer.bos_token_id]], dtype=np.int64)
33
34# run decoder model
35max_length = 200 # define maximum length of the sequence
36
37for _ in range(max_length):
38 decoder_outputs = decoder_session.run(
39 None,
40 {
41 "input_ids": decoder_input_ids, # input for the decoder
42 "encoder_hidden_states": encoder_hidden_states # outputs from the encoder
43 }
44 )
45
46 # extract logits and predict next token
47 logits = decoder_outputs[0]
48 predicted_token_id = np.argmax(logits[0, -1, :]) # get the predicted token ID from the logits
49
50 # if the predicted token is the EOS token, stop the generation
51 if predicted_token_id == tokenizer.eos_token_id:
52 break
53
54 # append predicted token ID to the decoder inputs for the next step
55 decoder_input_ids = np.concatenate([decoder_input_ids, np.array([[predicted_token_id]])], axis=-1)
56
57# decode the predicted token IDs into text
58predicted_text = tokenizer.decode(decoder_input_ids[0], skip_special_tokens=True)
59
60# print the generated caption
61print(predicted_text)