Views
No views yet
1from transformers import VisionEncoderDecoderModel, ViTFeatureExtractor, AutoTokenizer
2import torch
3from PIL import Image
4import pathlib
5import pandas as pd
6import numpy as np
7from IPython.core.display import HTML
8import os
9import requests
10
11class Image2Caption(object):
12 def __init__(self ,model_path = "nlpconnect/vit-gpt2-image-captioning",
13 device = torch.device("cuda" if torch.cuda.is_available() else "cpu"),
14 overwrite_encoder_checkpoint_path = None,
15 overwrite_token_model_path = None
16 ):
17 assert type(overwrite_token_model_path) == type("") or overwrite_token_model_path is None
18 assert type(overwrite_encoder_checkpoint_path) == type("") or overwrite_encoder_checkpoint_path is None
19 if overwrite_token_model_path is None:
20 overwrite_token_model_path = model_path
21 if overwrite_encoder_checkpoint_path is None:
22 overwrite_encoder_checkpoint_path = model_path
23 self.device = device
24 self.model = VisionEncoderDecoderModel.from_pretrained(model_path)
25 self.feature_extractor = ViTFeatureExtractor.from_pretrained(overwrite_encoder_checkpoint_path)
26 self.tokenizer = AutoTokenizer.from_pretrained(overwrite_token_model_path)
27 self.model = self.model.to(self.device)
28
29 def predict_to_df(self, image_paths):
30 img_caption_pred = self.predict_step(image_paths)
31 img_cation_df = pd.DataFrame(list(zip(image_paths, img_caption_pred)))
32 img_cation_df.columns = ["img", "caption"]
33 return img_cation_df
34 #img_cation_df.to_html(escape=False, formatters=dict(Country=path_to_image_html))
35
36 def predict_step(self ,image_paths, max_length = 128, num_beams = 4):
37 gen_kwargs = {"max_length": max_length, "num_beams": num_beams}
38 images = []
39 for image_path in image_paths:
40 #i_image = Image.open(image_path)
41 if image_path.startswith("http"):
42 i_image = Image.open(
43 requests.get(image_path, stream=True).raw
44 )
45 else:
46 i_image = Image.open(image_path)
47
48 if i_image.mode != "RGB":
49 i_image = i_image.convert(mode="RGB")
50 images.append(i_image)
51
52 pixel_values = self.feature_extractor(images=images, return_tensors="pt").pixel_values
53 pixel_values = pixel_values.to(self.device)
54
55 output_ids = self.model.generate(pixel_values, **gen_kwargs)
56
57 preds = self.tokenizer.batch_decode(output_ids, skip_special_tokens=True)
58 preds = [pred.strip() for pred in preds]
59 return preds
60
61def path_to_image_html(path):
62 return '<img src="'+ path + '" width="60" >'
63
64i2c_tiny_zh_obj = Image2Caption("svjack/vit-gpt-diffusion-zh",
65 overwrite_encoder_checkpoint_path = "google/vit-base-patch16-224",
66 overwrite_token_model_path = "IDEA-CCNL/Wenzhong-GPT2-110M"
67 )
68
69i2c_tiny_zh_obj.predict_step(
70 ["https://datasets-server.huggingface.co/assets/poloclub/diffusiondb/--/2m_all/train/28/image/image.jpg"]
71)
['"一个年轻男人的肖像,由Greg Rutkowski创作"。Artstation上的趋势"。"《刀锋战士》的艺术作品"。高度细节化。"电影般的灯光"。超现实主义。锐利的焦点。辛烷�']