1import os
2os.environ['LOWRES_RESIZE'] = '384x32'
3os.environ['HIGHRES_BASE'] = '0x32'
4os.environ['VIDEO_RESIZE'] = "0x64"
5os.environ['VIDEO_MAXRES'] = "480"
6os.environ['VIDEO_MINRES'] = "288"
7os.environ['MAXRES'] = '1536'
8os.environ['MINRES'] = '0'
9os.environ['REGIONAL_POOL'] = '2x'
10os.environ['FORCE_NO_DOWNSAMPLE'] = '1'
11os.environ['LOAD_VISION_EARLY'] = '1'
12os.environ['SKIP_LOAD_VIT'] = '1'
13
14
15import gradio as gr
16import torch
17import re
18from decord import VideoReader, cpu
19from PIL import Image
20import numpy as np
21import transformers
22import moviepy.editor as mp
23from typing import Dict, Optional, Sequence, List
24import librosa
25import whisper
26from ola.conversation import conv_templates, SeparatorStyle
27from ola.model.builder import load_pretrained_model
28from ola.utils import disable_torch_init
29from ola.datasets.preprocess import tokenizer_image_token, tokenizer_speech_image_token, tokenizer_speech_question_image_token
30from ola.mm_utils import get_model_name_from_path, KeywordsStoppingCriteria, process_anyres_video, process_anyres_highres_image_genli
31from ola.constants import IGNORE_INDEX, DEFAULT_IMAGE_TOKEN, IMAGE_TOKEN_INDEX, DEFAULT_SPEECH_TOKEN
32
33model_path = ""
34tokenizer, model, image_processor, _ = load_pretrained_model(model_path, None)
35model = model.to('cuda').eval()
36model = model.bfloat16()
37
38USE_SPEECH=False
39cur_dir = os.path.dirname(os.path.abspath(__file__))
40
41
42def load_audio(audio_file_name):
43 speech_wav, samplerate = librosa.load(audio_file_name, sr=16000)
44 if len(speech_wav.shape) > 1:
45 speech_wav = speech_wav[:, 0]
46 speech_wav = speech_wav.astype(np.float32)
47 CHUNK_LIM = 480000
48 SAMPLE_RATE = 16000
49 speechs = []
50 speech_wavs = []
51
52 if len(speech_wav) <= CHUNK_LIM:
53 speech = whisper.pad_or_trim(speech_wav)
54 speech_wav = whisper.pad_or_trim(speech_wav)
55 speechs.append(speech)
56 speech_wavs.append(torch.from_numpy(speech_wav).unsqueeze(0))
57 else:
58 for i in range(0, len(speech_wav), CHUNK_LIM):
59 chunk = speech_wav[i : i + CHUNK_LIM]
60 if len(chunk) < CHUNK_LIM:
61 chunk = whisper.pad_or_trim(chunk)
62 speechs.append(chunk)
63 speech_wavs.append(torch.from_numpy(chunk).unsqueeze(0))
64 mels = []
65 for chunk in speechs:
66 chunk = whisper.log_mel_spectrogram(chunk, n_mels=128).permute(1, 0).unsqueeze(0)
67 mels.append(chunk)
68
69 mels = torch.cat(mels, dim=0)
70 speech_wavs = torch.cat(speech_wavs, dim=0)
71 if mels.shape[0] > 25:
72 mels = mels[:25]
73 speech_wavs = speech_wavs[:25]
74
75 speech_length = torch.LongTensor([mels.shape[1]] * mels.shape[0])
76 speech_chunks = torch.LongTensor([mels.shape[0]])
77 return mels, speech_length, speech_chunks, speech_wavs
78
79def extract_audio(videos_file_path):
80 my_clip = mp.VideoFileClip(videos_file_path)
81 return my_clip.audio
82
83def ola_inference(multimodal, audio_path):
84 visual, text = multimodal["files"][0], multimodal["text"]
85 if visual.endswith("image2.png"):
86 modality = "video"
87 visual = f"{cur_dir}/case/case1.mp4"
88 if visual.endswith(".mp4"):
89 modality = "video"
90 else:
91 modality = "image"
92
93 # input audio and video, do not parse audio in the video, else parse audio in the video
94 if audio_path:
95 USE_SPEECH = True
96 elif modality == "video":
97 USE_SPEECH = True
98 else:
99 USE_SPEECH = False
100
101 speechs = []
102 speech_lengths = []
103 speech_wavs = []
104 speech_chunks = []
105 if modality == "video":
106 vr = VideoReader(visual, ctx=cpu(0))
107 total_frame_num = len(vr)
108 fps = round(vr.get_avg_fps())
109 uniform_sampled_frames = np.linspace(0, total_frame_num - 1, 64, dtype=int)
110 frame_idx = uniform_sampled_frames.tolist()
111 spare_frames = vr.get_batch(frame_idx).asnumpy()
112 video = [Image.fromarray(frame) for frame in spare_frames]
113 else:
114 image = [Image.open(visual)]
115 image_sizes = [image[0].size]
116
117 if USE_SPEECH and audio_path:
118 audio_path = audio_path
119 speech, speech_length, speech_chunk, speech_wav = load_audio(audio_path)
120 speechs.append(speech.bfloat16().to('cuda'))
121 speech_lengths.append(speech_length.to('cuda'))
122 speech_chunks.append(speech_chunk.to('cuda'))
123 speech_wavs.append(speech_wav.to('cuda'))
124 print('load audio')
125 elif USE_SPEECH and not audio_path:
126 # parse audio in the video
127 audio = extract_audio(visual)
128 audio.write_audiofile("./video_audio.wav")
129 video_audio_path = './video_audio.wav'
130 speech, speech_length, speech_chunk, speech_wav = load_audio(video_audio_path)
131 speechs.append(speech.bfloat16().to('cuda'))
132 speech_lengths.append(speech_length.to('cuda'))
133 speech_chunks.append(speech_chunk.to('cuda'))
134 speech_wavs.append(speech_wav.to('cuda'))
135 else:
136 speechs = [torch.zeros(1, 3000, 128).bfloat16().to('cuda')]
137 speech_lengths = [torch.LongTensor([3000]).to('cuda')]
138 speech_wavs = [torch.zeros([1, 480000]).to('cuda')]
139 speech_chunks = [torch.LongTensor([1]).to('cuda')]
140
141 conv_mode = "qwen_1_5"
142 if text:
143 qs = text
144 else:
145 qs = ''
146 if USE_SPEECH and audio_path:
147 qs = DEFAULT_IMAGE_TOKEN + "
148" + "User's question in speech: " + DEFAULT_SPEECH_TOKEN + '
149'
150 elif USE_SPEECH:
151 qs = DEFAULT_SPEECH_TOKEN + DEFAULT_IMAGE_TOKEN + "
152" + qs
153 else:
154 qs = DEFAULT_IMAGE_TOKEN + "
155" + qs
156
157 conv = conv_templates[conv_mode].copy()
158 conv.append_message(conv.roles[0], qs)
159 conv.append_message(conv.roles[1], None)
160 prompt = conv.get_prompt()
161 if USE_SPEECH and audio_path:
162 input_ids = tokenizer_speech_question_image_token(prompt, tokenizer, IMAGE_TOKEN_INDEX, return_tensors="pt").unsqueeze(0).to('cuda')
163 elif USE_SPEECH:
164 input_ids = tokenizer_speech_image_token(prompt, tokenizer, IMAGE_TOKEN_INDEX, return_tensors="pt").unsqueeze(0).to('cuda')
165 else:
166 input_ids = tokenizer_image_token(prompt, tokenizer, IMAGE_TOKEN_INDEX, return_tensors="pt").unsqueeze(0).to('cuda')
167
168 if modality == "video":
169 video_processed = []
170 for idx, frame in enumerate(video):
171 image_processor.do_resize = False
172 image_processor.do_center_crop = False
173 frame = process_anyres_video(frame, image_processor)
174
175 if frame_idx is not None and idx in frame_idx:
176 video_processed.append(frame.unsqueeze(0))
177 elif frame_idx is None:
178 video_processed.append(frame.unsqueeze(0))
179
180 if frame_idx is None:
181 frame_idx = np.arange(0, len(video_processed), dtype=int).tolist()
182
183 video_processed = torch.cat(video_processed, dim=0).bfloat16().to("cuda")
184 video_processed = (video_processed, video_processed)
185
186 video_data = (video_processed, (384, 384), "video")
187 else:
188 image_processor.do_resize = False
189 image_processor.do_center_crop = False
190 image_tensor, image_highres_tensor = [], []
191 for visual in image:
192 image_tensor_, image_highres_tensor_ = process_anyres_highres_image_genli(visual, image_processor)
193 image_tensor.append(image_tensor_)
194
195 image_highres_tensor.append(image_highres_tensor_)
196 if all(x.shape == image_tensor[0].shape for x in image_tensor):
197 image_tensor = torch.stack(image_tensor, dim=0)
198 if all(x.shape == image_highres_tensor[0].shape for x in image_highres_tensor):
199 image_highres_tensor = torch.stack(image_highres_tensor, dim=0)
200 if type(image_tensor) is list:
201 image_tensor = [_image.bfloat16().to("cuda") for _image in image_tensor]
202 else:
203 image_tensor = image_tensor.bfloat16().to("cuda")
204 if type(image_highres_tensor) is list:
205 image_highres_tensor = [_image.bfloat16().to("cuda") for _image in image_highres_tensor]
206 else:
207 image_highres_tensor = image_highres_tensor.bfloat16().to("cuda")
208
209 pad_token_ids = 151643
210
211 attention_masks = input_ids.ne(pad_token_ids).long().to('cuda')
212 stop_str = conv.sep if conv.sep_style != SeparatorStyle.TWO else conv.sep2
213 keywords = [stop_str]
214 stopping_criteria = KeywordsStoppingCriteria(keywords, tokenizer, input_ids)
215
216 gen_kwargs = {}
217
218 if "max_new_tokens" not in gen_kwargs:
219 gen_kwargs["max_new_tokens"] = 1024
220 if "temperature" not in gen_kwargs:
221 gen_kwargs["temperature"] = 0.2
222 if "top_p" not in gen_kwargs:
223 gen_kwargs["top_p"] = None
224 if "num_beams" not in gen_kwargs:
225 gen_kwargs["num_beams"] = 1
226
227 with torch.inference_mode():
228 if modality == "video":
229 output_ids = model.generate(
230 inputs=input_ids,
231 images=video_data[0][0],
232 images_highres=video_data[0][1],
233 modalities=video_data[2],
234 speech=speechs,
235 speech_lengths=speech_lengths,
236 speech_chunks=speech_chunks,
237 speech_wav=speech_wavs,
238 attention_mask=attention_masks,
239 use_cache=True,
240 stopping_criteria=[stopping_criteria],
241 do_sample=True if gen_kwargs["temperature"] > 0 else False,
242 temperature=gen_kwargs["temperature"],
243 top_p=gen_kwargs["top_p"],
244 num_beams=gen_kwargs["num_beams"],
245 max_new_tokens=gen_kwargs["max_new_tokens"],
246 )
247 else:
248 output_ids = model.generate(
249 inputs=input_ids,
250 images=image_tensor,
251 images_highres=image_highres_tensor,
252 image_sizes=image_sizes,
253 modalities=['image'],
254 speech=speechs,
255 speech_lengths=speech_lengths,
256 speech_chunks=speech_chunks,
257 speech_wav=speech_wavs,
258 attention_mask=attention_masks,
259 use_cache=True,
260 stopping_criteria=[stopping_criteria],
261 do_sample=True if gen_kwargs["temperature"] > 0 else False,
262 temperature=gen_kwargs["temperature"],
263 top_p=gen_kwargs["top_p"],
264 num_beams=gen_kwargs["num_beams"],
265 max_new_tokens=gen_kwargs["max_new_tokens"],
266 )
267
268 outputs = tokenizer.batch_decode(output_ids, skip_special_tokens=True)[0]
269 outputs = outputs.strip()
270 if outputs.endswith(stop_str):
271 outputs = outputs[:-len(stop_str)]
272 outputs = outputs.strip()
273 return outputs, None