Views
No views yet
1
2import torch
3import requests
4from PIL import Image
5from transformers import ViTFeatureExtractor, AutoTokenizer, VisionEncoderDecoderModel
6
7
8loc = "ydshieh/vit-gpt2-coco-en"
9
10feature_extractor = ViTFeatureExtractor.from_pretrained(loc)
11tokenizer = AutoTokenizer.from_pretrained(loc)
12model = VisionEncoderDecoderModel.from_pretrained(loc)
13model.eval()
14
15
16def predict(image):
17
18 pixel_values = feature_extractor(images=image, return_tensors="pt").pixel_values
19
20 with torch.no_grad():
21 output_ids = model.generate(pixel_values, max_length=16, num_beams=4, return_dict_in_generate=True).sequences
22
23 preds = tokenizer.batch_decode(output_ids, skip_special_tokens=True)
24 preds = [pred.strip() for pred in preds]
25
26 return preds
27
28
29# We will verify our results on an image of cute cats
30url = "http://images.cocodataset.org/val2017/000000039769.jpg"
31with Image.open(requests.get(url, stream=True).raw) as image:
32 preds = predict(image)
33
34print(preds)
35# should produce
36# ['a cat laying on top of a couch next to another cat']
371
2import jax
3import requests
4from PIL import Image
5from transformers import ViTFeatureExtractor, AutoTokenizer, FlaxVisionEncoderDecoderModel
6
7
8loc = "ydshieh/vit-gpt2-coco-en"
9
10feature_extractor = ViTFeatureExtractor.from_pretrained(loc)
11tokenizer = AutoTokenizer.from_pretrained(loc)
12model = FlaxVisionEncoderDecoderModel.from_pretrained(loc)
13
14gen_kwargs = {"max_length": 16, "num_beams": 4}
15
16
17# This takes sometime when compiling the first time, but the subsequent inference will be much faster
18@jax.jit
19def generate(pixel_values):
20 output_ids = model.generate(pixel_values, **gen_kwargs).sequences
21 return output_ids
22
23
24def predict(image):
25
26 pixel_values = feature_extractor(images=image, return_tensors="np").pixel_values
27 output_ids = generate(pixel_values)
28 preds = tokenizer.batch_decode(output_ids, skip_special_tokens=True)
29 preds = [pred.strip() for pred in preds]
30
31 return preds
32
33
34# We will verify our results on an image of cute cats
35url = "http://images.cocodataset.org/val2017/000000039769.jpg"
36with Image.open(requests.get(url, stream=True).raw) as image:
37 preds = predict(image)
38
39print(preds)
40# should produce
41# ['a cat laying on top of a couch next to another cat']
42