Views
No views yet
1import mimetypes
2import os
3from io import BytesIO
4from typing import Union
5import cv2
6import requests
7import torch
8import transformers
9from PIL import Image
10from torchvision.transforms import Compose, Resize, ToTensor
11from tqdm import tqdm
12import sys
13
14from otter.modeling_otter import OtterForConditionalGeneration
15
16
17# Disable warnings
18requests.packages.urllib3.disable_warnings()
19
20# ------------------- Utility Functions -------------------
21
22
23def get_content_type(file_path):
24 content_type, _ = mimetypes.guess_type(file_path)
25 return content_type
26
27
28# ------------------- Image and Video Handling Functions -------------------
29
30def get_image(url: str) -> Union[Image.Image, list]:
31 if "://" not in url: # Local file
32 content_type = get_content_type(url)
33 else: # Remote URL
34 content_type = requests.head(url, stream=True, verify=False).headers.get("Content-Type")
35
36 if "image" in content_type:
37 if "://" not in url: # Local file
38 return Image.open(url)
39 else: # Remote URL
40 return Image.open(requests.get(url, stream=True, verify=False).raw)
41 else:
42 raise ValueError("Invalid content type. Expected image or video.")
43
44
45# ------------------- OTTER Prompt and Response Functions -------------------
46
47
48def get_formatted_prompt(prompt: str, in_context_prompts: list = []) -> str:
49 in_context_string = ""
50 for in_context_prompt, in_context_answer in in_context_prompts:
51 in_context_string += f"<image>User: {in_context_prompt} GPT:<answer> {in_context_answer}<|endofchunk|>"
52 return f"{in_context_string}<image>User: {prompt} GPT:<answer>"
53
54
55def get_response(image_list, prompt: str, model=None, image_processor=None, in_context_prompts: list = []) -> str:
56 input_data = image_list
57
58 if isinstance(input_data, Image.Image):
59 vision_x = image_processor.preprocess([input_data], return_tensors="pt")["pixel_values"].unsqueeze(1).unsqueeze(0)
60 elif isinstance(input_data, list): # list of video frames
61 vision_x = image_processor.preprocess(input_data, return_tensors="pt")["pixel_values"].unsqueeze(1).unsqueeze(0)
62 else:
63 raise ValueError("Invalid input data. Expected PIL Image or list of video frames.")
64
65 lang_x = model.text_tokenizer(
66 [
67 get_formatted_prompt(prompt, in_context_prompts),
68 ],
69 return_tensors="pt",
70 )
71 bad_words_id = tokenizer(["User:", "GPT1:", "GFT:", "GPT:"], add_special_tokens=False).input_ids
72 generated_text = model.generate(
73 vision_x=vision_x.to(model.device),
74 lang_x=lang_x["input_ids"].to(model.device),
75 attention_mask=lang_x["attention_mask"].to(model.device),
76 max_new_tokens=512,
77 num_beams=3,
78 no_repeat_ngram_size=3,
79 bad_words_ids=bad_words_id,
80 )
81 parsed_output = (
82 model.text_tokenizer.decode(generated_text[0])
83 .split("<answer>")[-1]
84 .lstrip()
85 .rstrip()
86 .split("<|endofchunk|>")[0]
87 .lstrip()
88 .rstrip()
89 .lstrip('"')
90 .rstrip('"')
91 )
92 return parsed_output
93
94
95# ------------------- Main Function -------------------
96
97if __name__ == "__main__":
98 model = OtterForConditionalGeneration.from_pretrained("luodian/OTTER-9B-LA-InContext", device_map="auto")
99 model.text_tokenizer.padding_side = "left"
100 tokenizer = model.text_tokenizer
101 image_processor = transformers.CLIPImageProcessor()
102 model.eval()
103
104 while True:
105 urls = [
106 "https://images.cocodataset.org/train2017/000000339543.jpg",
107 "https://images.cocodataset.org/train2017/000000140285.jpg",
108 ]
109
110 encoded_frames_list = []
111 for url in urls:
112 frames = get_image(url)
113 encoded_frames_list.append(frames)
114
115 in_context_prompts = []
116 in_context_examples = [
117 "What does the image describe?::A family is taking picture in front of a snow mountain.",
118 ]
119 for in_context_input in in_context_examples:
120 in_context_prompt, in_context_answer = in_context_input.split("::")
121 in_context_prompts.append((in_context_prompt.strip(), in_context_answer.strip()))
122
123 # prompts_input = input("Enter the prompts separated by commas (or type 'quit' to exit): ")
124 prompts_input = "What does the image describe?"
125
126 prompts = [prompt.strip() for prompt in prompts_input.split(",")]
127
128 for prompt in prompts:
129 print(f"\nPrompt: {prompt}")
130 response = get_response(encoded_frames_list, prompt, model, image_processor, in_context_prompts)
131 print(f"Response: {response}")
132
133 if prompts_input.lower() == "quit":
134 break