diff --git a/easyanimate/video_caption/README.md b/easyanimate/video_caption/README.md new file mode 100644 index 0000000..00ad98b --- /dev/null +++ b/easyanimate/video_caption/README.md @@ -0,0 +1,59 @@ +# Video Caption +EasyAnimate uses multi-modal LLMs to generate captions for frames extracted from the video firstly, and then employs LLMs to summarize and refine the generated frame captions into the final video caption. By leveraging [sglang](https://github.com/sgl-project/sglang)/[vLLM](https://github.com/vllm-project/vllm) and [accelerate distributed inference](https://huggingface.co/docs/accelerate/en/usage_guides/distributed_inference), the entire processing could be very fast. + +## Quick Start +1. Cloud usage: AliyunDSW/Docker + + Check [README.md](../README.md) for details. + +2. Local usage + + ```shell + # Install EasyAnimate requirements firstly. + cd EasyAnimate && pip install -r requirements.txt + + # Install additional requirements for video caption. + cd easyanimate/video_caption && pip install -r requirements.txt + ``` + +## How to use + +1. Prepare videos. + + The input for video caption can be a video folder or a metadata file (txt/csv/jsonl) containing the video path column. Please check `get_video_path_list` function in [utils/video_utils.py](utils/video_utils.py) for details. + +2. Generate frame captions. + + We have conducted a detailed and manual comparison of open sourced multi-modal LLMs such as [Qwen-VL](https://huggingface.co/Qwen/Qwen-VL), [ShareGPT4V-7B](https://huggingface.co/Lin-Chen/ShareGPT4V-7B), [deepseek-vl-7b-chat](https://huggingface.co/deepseek-ai/deepseek-vl-7b-chat) and etc. And we found that [llava-v1.6-vicuna-7b](https://huggingface.co/liuhaotian/llava-v1.6-vicuna-7b) is capable of generating more detailed captions with fewer hallucinations. Additionally, it is supported by serving engines like [sglang](https://github.com/sgl-project/sglang) and [lmdepoly](https://github.com/InternLM/lmdeploy), enabling faster inference. + + ```shell + CUDA_VISIBLE_DEVICES=0 python caption_video_frame.py \ + --video_folder="your-video-folder/" + --frame_sample_method="extract_mid_frame" \ + --num_sampled_frames=1 \ + --image_caption_model_name="llava-v1.6-vicuna-7b" \ + --image_caption_prompt="Please describe this image in detail." \ + --saved_path="video_frame_caption.jsonl" + ``` + + If you cannot access to Huggingface, you can run `export HF_ENDPOINT=https://hf-mirror.com` before the above command to download the image caption model automatically. + +3. Summary frame captions. + + ```shell + CUDA_VISIBLE_DEVICES=0 python caption_summary.py \ + --video_metadata_path="video_frame_caption_result.jsonl" \ + --video_path_column="video_path" \ + --caption_column="sampled_frame_caption" \ + --summary_model_name="mistralai/Mistral-7B-Instruct-v0.2" \ + --summary_prompt="You are a helpful video description generator. I'll give you a description of the middle frame of the video clip, \ + which you need to summarize it into a description of the video clip. \ + Please provide your video description following these requirements: \ + 1. Describe the basic and necessary information of the video in the third person, be as concise as possible. \ + 2. Output the video description directly. Begin with 'In this video'. \ + 3. Limit the video description within 100 words. \ + Here is the mid-frame description: " \ + --saved_path="video_summary_caption.jsonl" + ``` + + If you cannot access to Huggingface, you can run `export HF_ENDPOINT=https://hf-mirror.com` before the above command to download the summary caption model automatically. \ No newline at end of file diff --git a/easyanimate/video_caption/caption_summary.py b/easyanimate/video_caption/caption_summary.py new file mode 100644 index 0000000..d0c99f4 --- /dev/null +++ b/easyanimate/video_caption/caption_summary.py @@ -0,0 +1,134 @@ +import argparse +import os +import re +from tqdm import tqdm + +import pandas as pd +from vllm import LLM, SamplingParams + +from utils.logger import logger + + +def parse_args(): + parser = argparse.ArgumentParser(description="Recaption the video frame.") + parser.add_argument( + "--video_metadata_path", type=str, required=True, help="The path to the video dataset metadata (csv/jsonl)." + ) + parser.add_argument( + "--video_path_column", + type=str, + default="video_path", + help="The column contains the video path (an absolute path or a relative path w.r.t the video_folder).", + ) + parser.add_argument( + "--caption_column", + type=str, + default="sampled_frame_caption", + help="The column contains the sampled_frame_caption.", + ) + parser.add_argument( + "--remove_quotes", + action="store_true", + help="Whether to remove quotes from caption.", + ) + parser.add_argument( + "--batch_size", + type=int, + default=10, + required=False, + help="The batch size for the video caption.", + ) + parser.add_argument( + "--summary_model_name", + type=str, + default="mistralai/Mistral-7B-Instruct-v0.2", + ) + parser.add_argument( + "--summary_prompt", + type=str, + default=( + "You are a helpful video description generator. I'll give you a description of the middle frame of the video clip, " + "which you need to summarize it into a description of the video clip." + "Please provide your video description following these requirements: " + "1. Describe the basic and necessary information of the video in the third person, be as concise as possible. " + "2. Output the video description directly. Begin with 'In this video'. " + "3. Limit the video description within 100 words. " + "Here is the mid-frame description: " + ), + ) + parser.add_argument("--saved_path", type=str, required=True, help="The save path to the output results (csv/jsonl).") + parser.add_argument("--saved_freq", type=int, default=1000, help="The frequency to save the output results.") + + args = parser.parse_args() + return args + + +def main(): + args = parse_args() + + if args.video_metadata_path.endswith(".csv"): + video_metadata_df = pd.read_csv(args.video_metadata_path) + elif args.video_metadata_path.endswith(".jsonl"): + video_metadata_df = pd.read_json(args.video_metadata_path, lines=True) + else: + raise ValueError("The video_metadata_path must end with .csv or .jsonl.") + video_path_list = video_metadata_df[args.video_path_column].tolist() + sampled_frame_caption_list = video_metadata_df[args.caption_column].tolist() + + if not (args.saved_path.endswith(".csv") or args.saved_path.endswith(".jsonl")): + raise ValueError("The saved_path must end with .csv or .jsonl.") + + if os.path.exists(args.saved_path): + if args.saved_path.endswith(".csv"): + saved_metadata_df = pd.read_csv(args.saved_path) + elif args.saved_path.endswith(".jsonl"): + saved_metadata_df = pd.read_json(args.saved_path, lines=True) + saved_video_path_list = saved_metadata_df[args.video_path_column].tolist() + video_path_list = list(set(video_path_list) - set(saved_video_path_list)) + video_metadata_df.set_index(args.video_path_column, inplace=True) + video_metadata_df = video_metadata_df.loc[video_path_list] + sampled_frame_caption_list = video_metadata_df[args.caption_column].tolist() + logger.info(f"Resume from {args.saved_path}: {len(saved_video_path_list)} processed and {len(video_path_list)} to be processed.") + + sampling_params = SamplingParams(temperature=0.8, top_p=0.95, max_tokens=256) + summary_model = LLM(model=args.summary_model_name, trust_remote_code=True) + + result_dict = {"video_path": [], "summary_model": [], "summary_caption": []} + + for i in tqdm(range(0, len(sampled_frame_caption_list), args.batch_size)): + batch_video_path = video_path_list[i: i + args.batch_size] + batch_caption = sampled_frame_caption_list[i : i + args.batch_size] + batch_prompt = [] + for caption in batch_caption: + if args.remove_quotes: + caption = re.sub(r'(["\']).*?\1', "", caption) + batch_prompt.append("user:" + args.summary_prompt + str(caption) + "\n assistant:") + batch_output = summary_model.generate(batch_prompt, sampling_params) + + result_dict["video_path"].extend(batch_video_path) + result_dict["summary_model"].extend([args.summary_model_name] * len(batch_caption)) + result_dict["summary_caption"].extend([output.outputs[0].text.rstrip() for output in batch_output]) + + # Save the metadata every args.saved_freq. + if i != 0 and ((i // args.batch_size) % args.saved_freq) == 0: + result_df = pd.DataFrame(result_dict) + if args.saved_path.endswith(".csv"): + header = True if not os.path.exists(args.saved_path) else False + result_df.to_csv(args.saved_path, header=header, index=False, mode="a") + elif args.saved_path.endswith(".jsonl"): + result_df.to_json(args.saved_path, orient="records", lines=True, mode="a") + logger.info(f"Save result to {args.saved_path}.") + + result_dict = {"video_path": [], "summary_model": [], "summary_caption": []} + + result_df = pd.DataFrame(result_dict) + if args.saved_path.endswith(".csv"): + header = True if not os.path.exists(args.saved_path) else False + result_df.to_csv(args.saved_path, header=header, index=False, mode="a") + elif args.saved_path.endswith(".jsonl"): + result_df.to_json(args.saved_path, orient="records", lines=True, mode="a") + logger.info(f"Save the final result to {args.saved_path}.") + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/easyanimate/video_caption/caption_video_frame.py b/easyanimate/video_caption/caption_video_frame.py new file mode 100644 index 0000000..f1481aa --- /dev/null +++ b/easyanimate/video_caption/caption_video_frame.py @@ -0,0 +1,256 @@ +import argparse +import copy +import os + +import pandas as pd +from accelerate import PartialState +from accelerate.utils import gather_object +from tqdm import tqdm +from torch.utils.data import DataLoader + +from utils.image_captioner import QwenVLChat, InternLMXComposer2, LLaVASRT +from utils.logger import logger +from utils.video_dataset import VideoDataset, collate_fn +from utils.video_utils import get_video_path_list, extract_frames + + +ACCELERATE_SUPPORTED_MODELS = ["Qwen-VL-Chat", "internlm-xcomposer2-vl-7b"] +SGLANG_SUPPORTED_MODELS = ["llava-v1.6-vicuna-7b"] + + +def parse_args(): + parser = argparse.ArgumentParser(description="Recaption the video frame.") + parser.add_argument("--video_folder", type=str, default="", help="The video folder.") + parser.add_argument( + "--video_metadata_path", type=str, default=None, help="The path to the video dataset metadata (csv/jsonl/txt)." + ) + parser.add_argument( + "--video_path_column", + type=str, + default="video_path", + help="The column contains the video path (an absolute path or a relative path w.r.t the video_folder).", + ) + parser.add_argument( + "--batch_size", + type=int, + default=10, + required=False, + help="The batch size for the video dataset.", + ) + parser.add_argument( + "--frame_sample_method", + type=str, + choices=["mid", "uniform"], + default="mid", + ) + parser.add_argument( + "--num_sampled_frames", + type=int, + default=1, + help="num_sampled_frames", + ) + parser.add_argument( + "--image_caption_model_name", + type=str, + choices=ACCELERATE_SUPPORTED_MODELS + SGLANG_SUPPORTED_MODELS, + default="internlm-xcomposer2-vl-7b", + ) + parser.add_argument( + "--image_caption_model_quantized", type=bool, default=True, help="Whether to use the quantized image caption model." + ) + parser.add_argument( + "--image_caption_prompt", + type=str, + default="Describe this image and its style in a very detailed manner.", + ) + parser.add_argument( + "--output_dir", + type=str, + required=True, + help="The directory to create the subfolder (named with the video name) to indicate the video has been processed.", + ) + parser.add_argument("--saved_path", type=str, required=True, help="The save path to the output results (csv/jsonl).") + parser.add_argument("--saved_freq", type=int, default=1000, help="The frequency to save the output results.") + + args = parser.parse_args() + return args + + +def accelerate_inference(args, video_path_list): + state = PartialState() + device = state.device + if state.num_processes == 1: + device = "cuda:0" + if args.image_caption_model_name == "internlm-xcomposer2-vl-7b": + image_caption_model = InternLMXComposer2(device=device, quantized=args.image_caption_model_quantized) + elif args.image_caption_model_name == "Qwen-VL-Chat": + image_caption_model = QwenVLChat(device=device, quantized=args.image_caption_model_quantized) + + if state.is_main_process: + os.makedirs(args.output_dir, exist_ok=True) + result_list = [] + with state.split_between_processes(video_path_list) as splitted_video_path_list: + for i, video_path in enumerate(tqdm(splitted_video_path_list, desc=f"{state.device}")): + video_id = os.path.splitext(os.path.basename(video_path))[0] + try: + if not os.path.exists(video_path): + print(f"Video {video_id} does not exist. Pass it.") + continue + sampled_frame_list, sampled_frame_idx_list = extract_frames(video_path, num_sample_frames=args.num_sample_frames) + except Exception as e: + print(f"Failed to extract frames from video {video_id}. Error is {e}.") + + video_recaption_output_dir = os.path.join(args.output_dir, video_id) + if os.path.exists(video_recaption_output_dir): + print(f"Video {video_id} has been processed. Pass it.") + continue + else: + os.makedirs(video_recaption_output_dir) + + caption_list = [] + for frame, frame_idx in zip(sampled_frame_list, sampled_frame_idx_list): + frame_path = f"{args.output_dir}/{video_id}_{frame_idx}.png" + frame.save(frame_path) + try: + response, _ = image_caption_model(args.image_caption_prompt, frame_path) + except Exception as e: + print(f"Failed to caption video {video_id}. Error is {e}.") + finally: + os.remove(frame_path) + caption_list.append(response) + + result_meta = {} + if args.video_folder == "": + result_meta[args.video_path_column] = video_path + else: + result_meta[args.video_path_column] = os.path.basename(video_path) + result_meta["image_caption_model"] = args.image_caption_model_name + result_meta["prompt"] = args.image_caption_prompt + result_meta["sampled_frame_idx"] = sampled_frame_idx_list + result_meta["sampled_frame_caption"] = caption_list + result_list.append(copy.deepcopy(result_meta)) + + # Save the metadata in the main process. + if i != 0 and i % args.saved_freq == 0: + state.wait_for_everyone() + gathered_result_list = gather_object(result_list) + if state.is_main_process: + result_df = pd.DataFrame(gathered_result_list) + if args.saved_path.endswith(".csv"): + result_df.to_csv(args.saved_path, index=False) + elif args.saved_path.endswith(".jsonl"): + result_df.to_json(args.saved_path, orient="records", lines=True) + print(f"Save result to {args.saved_path}.") + + # Wait for all processes to finish and gather the final result. + state.wait_for_everyone() + gathered_result_list = gather_object(result_list) + # Save the metadata in the main process. + if state.is_main_process: + result_df = pd.DataFrame(gathered_result_list) + if args.saved_path.endswith(".csv"): + result_df.to_csv(args.saved_path, index=False) + elif args.saved_path.endswith(".jsonl"): + result_df.to_json(args.saved_path, orient="records", lines=True) + print(f"Save the final result to {args.saved_path}.") + + +def sglang_inference(args, video_path_list): + if args.image_caption_model_name == "llava-v1.6-vicuna-7b": + image_caption_model = LLaVASRT() + + result_dict = { + "video_path": [], + "image_caption_model": [], + "prompt": [], + 'sampled_frame_idx': [], + "sampled_frame_caption": [] + } + + video_dataset = VideoDataset( + video_path_list=video_path_list, + sample_method=args.frame_sample_method, + num_sampled_frames=args.num_sampled_frames + ) + video_loader = DataLoader(video_dataset, batch_size=args.batch_size, num_workers=16, collate_fn=collate_fn) + for idx, batch in enumerate(tqdm(video_loader)): + if len(batch) == 0: + continue + batch_video_path, batch_frame_idx = batch["video_path"], batch["sampled_frame_idx"] + # [batch_size, num_sampled_frames, H, W, C] => [batch_size * num_sampled_frames, H, W, C]. + batch_frame = [] + for item_sampled_frame in batch["sampled_frame"]: + batch_frame.extend([frame for frame in item_sampled_frame]) + + try: + response_list, _ = image_caption_model([args.image_caption_prompt] * len(batch_frame), batch_frame) + response_list = [response_list[i:i + args.num_sampled_frames] for i in range(0, len(response_list), args.num_sampled_frames)] + except Exception as e: + logger.error(f"Failed to caption video {batch_video_path}. Error is {e}.") + + result_dict["video_path"].extend(batch_video_path) + result_dict["image_caption_model"].extend([args.image_caption_model_name] * len(batch_video_path)) + result_dict["prompt"].extend([args.image_caption_prompt] * len(batch_video_path)) + result_dict["sampled_frame_idx"].extend(batch_frame_idx) + result_dict["sampled_frame_caption"].extend(response_list) + + # Save the metadata in the main process. + if idx != 0 and idx % args.saved_freq == 0: + result_df = pd.DataFrame(result_dict) + if args.saved_path.endswith(".csv"): + header = True if not os.path.exists(args.saved_path) else False + result_df.to_csv(args.saved_path, header=header, index=False, mode="a") + elif args.saved_path.endswith(".jsonl"): + result_df.to_json(args.saved_path, orient="records", lines=True, mode="a") + logger.info(f"Save result to {args.saved_path}.") + + result_dict = { + "video_path": [], + "image_caption_model": [], + "prompt": [], + 'sampled_frame_idx': [], + "sampled_frame_caption": [] + } + + if len(result_dict["video_path"]) != 0: + result_df = pd.DataFrame(result_dict) + if args.saved_path.endswith(".csv"): + header = True if not os.path.exists(args.saved_path) else False + result_df.to_csv(args.saved_path, header=header, index=False, mode="a") + elif args.saved_path.endswith(".jsonl"): + result_df.to_json(args.saved_path, orient="records", lines=True, mode="a") + logger.info(f"Save the final result to {args.saved_path}.") + + +def main(): + args = parse_args() + + video_path_list = get_video_path_list( + video_folder=args.video_folder, + video_metadata_path=args.video_metadata_path, + video_path_column=args.video_path_column + ) + + if not (args.saved_path.endswith(".csv") or args.saved_path.endswith(".jsonl")): + raise ValueError("The saved_path must end with .csv or .jsonl.") + + if os.path.exists(args.saved_path): + if args.saved_path.endswith(".csv"): + saved_metadata_df = pd.read_csv(args.saved_path) + elif args.saved_path.endswith(".jsonl"): + saved_metadata_df = pd.read_json(args.saved_path, lines=True) + saved_video_path_list = saved_metadata_df[args.video_path_column].tolist() + saved_video_path_list = [os.path.join(args.video_folder, path) for path in saved_video_path_list] + video_path_list = list(set(video_path_list) - set(saved_video_path_list)) + logger.info(f"Resume from {args.saved_path}: {len(saved_video_path_list)} processed and {len(video_path_list)} to be processed.") + + if args.image_caption_model_name in SGLANG_SUPPORTED_MODELS: + sglang_inference(args, video_path_list) + elif args.image_caption_model_name in ACCELERATE_SUPPORTED_MODELS: + accelerate_inference(args, video_path_list) + else: + raise ValueError(f"The {args.image_caption_model_name} is not supported.") + + +if __name__ == "__main__": + main() diff --git a/easyanimate/video_caption/requirements.txt b/easyanimate/video_caption/requirements.txt new file mode 100644 index 0000000..a453251 --- /dev/null +++ b/easyanimate/video_caption/requirements.txt @@ -0,0 +1,5 @@ +pandas>=2.0.0 +auto_gptq +vllm +sglang[srt] +func_timeout \ No newline at end of file diff --git a/easyanimate/video_caption/utils/image_captioner.py b/easyanimate/video_caption/utils/image_captioner.py new file mode 100644 index 0000000..0303288 --- /dev/null +++ b/easyanimate/video_caption/utils/image_captioner.py @@ -0,0 +1,168 @@ +import os +import time +from datetime import datetime +from pathlib import Path +from typing import List, Tuple, Union + +import auto_gptq +import sglang as sgl +import torch +from auto_gptq.modeling import BaseGPTQForCausalLM +from PIL import Image +from transformers import AutoModelForCausalLM, AutoTokenizer + +from utils.logger import logger + +TMP_DIR = "./tmp" + + +def get_timestamp(): + timestamp_ns = int(time.time_ns()) + milliseconds = timestamp_ns // 1000000 + formatted_time = datetime.fromtimestamp(milliseconds / 1000).strftime("%Y-%m-%d_%H-%M-%S-%f")[:-3] + + return formatted_time + + +class QwenVLChat: + def __init__(self, device: str = "cuda:0", quantized: bool = False) -> None: + if quantized: + self.model = AutoModelForCausalLM.from_pretrained( + "Qwen/Qwen-VL-Chat-Int4", device_map=device, trust_remote_code=True + ).eval() + self.tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen-VL-Chat-Int4", trust_remote_code=True) + else: + self.model = AutoModelForCausalLM.from_pretrained( + "Qwen/Qwen-VL-Chat", device_map=device, trust_remote_code=True, fp16=True + ).eval() + self.tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen-VL-Chat", trust_remote_code=True) + + def __call__(self, prompt: str, image: str) -> Tuple[str, str]: + query = self.tokenizer.from_list_format([{"image": image}, {"text": prompt}]) + response, history = self.model.chat(self.tokenizer, query=query, history=[]) + + return response, history + + +class InternLMXComposer2QForCausalLM(BaseGPTQForCausalLM): + layers_block_name = "model.layers" + outside_layer_modules = [ + "vit", + "vision_proj", + "model.tok_embeddings", + "model.norm", + "output", + ] + inside_layer_modules = [ + ["attention.wqkv.linear"], + ["attention.wo.linear"], + ["feed_forward.w1.linear", "feed_forward.w3.linear"], + ["feed_forward.w2.linear"], + ] + + +class InternLMXComposer2: + def __init__(self, device: str = "cuda:0", quantized: bool = True): + if quantized: + auto_gptq.modeling._base.SUPPORTED_MODELS = ["internlm"] + self.model = InternLMXComposer2QForCausalLM.from_quantized( + "internlm/internlm-xcomposer2-vl-7b-4bit", trust_remote_code=True, device=device + ).eval() + self.tokenizer = AutoTokenizer.from_pretrained("internlm/internlm-xcomposer2-vl-7b-4bit", trust_remote_code=True) + else: + # Setting fp16=True does not work. See https://huggingface.co/internlm/internlm-xcomposer2-vl-7b/discussions/1. + self.model = ( + AutoModelForCausalLM.from_pretrained( + "internlm/internlm-xcomposer2-vl-7b", device_map=device, trust_remote_code=True + ) + .eval() + .to(torch.float16) + ) + self.tokenizer = AutoTokenizer.from_pretrained("internlm/internlm-xcomposer2-vl-7b", trust_remote_code=True) + + def __call__(self, prompt: str, image: str): + if not prompt.startswith(""): + prompt = "" + prompt + with torch.cuda.amp.autocast(), torch.no_grad(): + response, history = self.model.chat(self.tokenizer, query=prompt, image=image, history=[], do_sample=False) + return response, history + + +class LLaVASRT: + def __init__(self, device: str = "cuda:0", quantized: bool = True): + runtime = sgl.Runtime(model_path="liuhaotian/llava-v1.6-vicuna-7b", tokenizer_path="llava-hf/llava-1.5-7b-hf") + sgl.set_default_backend(runtime) + logger.info( + f"Start the SGLang runtime for llava-v1.6-vicuna-7b with chat template: {runtime.endpoint.chat_template.name}. " + "Input parameter device and quantized do not take effect." + ) + if not os.path.exists(TMP_DIR): + os.makedirs(TMP_DIR, exist_ok=True) + + @sgl.function + def image_qa(s, prompt: str, image: str): + s += sgl.user(sgl.image(image) + prompt) + s += sgl.assistant(sgl.gen("answer")) + + def __call__(self, prompt: Union[str, List[str]], image: Union[str, Image.Image, List[str]]): + pil_input_flag = False + if isinstance(prompt, str) and (isinstance(image, str) or isinstance(image, Image.Image)): + if isinstance(image, Image.Image): + pil_input_flag = True + image_path = os.path.join(TMP_DIR, get_timestamp() + ".jpg") + image.save(image_path) + state = self.image_qa.run(prompt=prompt, image=image, max_new_tokens=256) + # Post-process. + if pil_input_flag: + os.remove(image) + + return state["answer"], state + elif isinstance(prompt, list) and isinstance(image, list): + assert len(prompt) == len(image) + if isinstance(image[0], Image.Image): + pil_input_flag = True + image_path = [os.path.join(TMP_DIR, get_timestamp() + f"-{i}" + ".jpg") for i in range(len(image))] + for i in range(len(image)): + image[i].save(image_path[i]) + image = image_path + batch_query = [{"prompt": p, "image": img} for p, img in zip(prompt, image)] + state = self.image_qa.run_batch(batch_query, max_new_tokens=256) + # Post-process. + if pil_input_flag: + for i in range(len(image)): + os.remove(image[i]) + + return [s["answer"] for s in state], state + else: + raise ValueError("Input prompt and image must be both strings or list of strings with the same length.") + + +if __name__ == "__main__": + image_folder = "demo/" + wildcard_list = ["*.jpg", "*.png"] + image_list = [] + for wildcard in wildcard_list: + image_list.extend([str(image_path) for image_path in Path(image_folder).glob(wildcard)]) + qwen_vl_chat = QwenVLChat(device="cuda:0", quantized=True) + qwen_vl_prompt = "Please describe this image in detail." + for image in image_list: + response, _ = qwen_vl_chat(qwen_vl_prompt, image) + print(image, response) + + internlm2_vl = InternLMXComposer2(device="cuda:0", quantized=False) + internlm2_vl_prompt = "Please describe this image in detail." + for image in image_list: + response, _ = internlm2_vl(internlm2_vl_prompt, image) + print(image, response) + + # # SGLang need the exclusive GPU and cannot re-initialize CUDA in forked subprocess. + # llava_srt = LLaVASRT() + # # Batch inference. + # llava_srt_prompt = ["Please describe this image in detail."] * len(image_list) + # response, _ = llava_srt(llava_srt_prompt, image_list) + # print(response) + # Single inference. + # llava_srt_prompt = "Please describe this image in detail." + # for image in image_list: + # response, _ = llava_srt(llava_srt_prompt, image) + # print(image, response) diff --git a/easyanimate/video_caption/utils/logger.py b/easyanimate/video_caption/utils/logger.py new file mode 100644 index 0000000..754eaf6 --- /dev/null +++ b/easyanimate/video_caption/utils/logger.py @@ -0,0 +1,36 @@ +# Borrowed from sd-webui-controlnet/scripts/logging.py +import copy +import logging +import sys + + +class ColoredFormatter(logging.Formatter): + COLORS = { + "DEBUG": "\033[0;36m", # CYAN + "INFO": "\033[0;32m", # GREEN + "WARNING": "\033[0;33m", # YELLOW + "ERROR": "\033[0;31m", # RED + "CRITICAL": "\033[0;37;41m", # WHITE ON RED + "RESET": "\033[0m", # RESET COLOR + } + + def format(self, record): + colored_record = copy.copy(record) + levelname = colored_record.levelname + seq = self.COLORS.get(levelname, self.COLORS["RESET"]) + colored_record.levelname = f"{seq}{levelname}{self.COLORS['RESET']}" + return super().format(colored_record) + + +# Create a new logger +logger = logging.getLogger("VideoCaption") +logger.propagate = False + +# Add handler if we don't have one. +if not logger.handlers: + handler = logging.StreamHandler(sys.stdout) + handler.setFormatter(ColoredFormatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s")) + logger.addHandler(handler) + +# Configure logger +logger.setLevel("INFO") diff --git a/easyanimate/video_caption/utils/video_dataset.py b/easyanimate/video_caption/utils/video_dataset.py new file mode 100644 index 0000000..537c411 --- /dev/null +++ b/easyanimate/video_caption/utils/video_dataset.py @@ -0,0 +1,83 @@ +from pathlib import Path + +import pandas as pd +from func_timeout import FunctionTimedOut, func_timeout +from torch.utils.data import DataLoader, Dataset + +from utils.logger import logger +from utils.video_utils import get_video_path_list, extract_frames + +ALL_VIDEO_EXT = set(["mp4", "webm", "mkv", "avi", "flv", "mov"]) +VIDEO_READER_TIMEOUT = 10 + + +def collate_fn(batch): + batch = list(filter(lambda x: x is not None, batch)) + if len(batch) != 0: + return {k: [item[k] for item in batch] for k in batch[0].keys()} + return {} + + +class VideoDataset(Dataset): + def __init__( + self, + video_path_list=None, + video_folder=None, + video_metadata_path=None, + video_path_column=None, + sample_method="mid", + num_sampled_frames=1, + num_sample_stride=None, + ): + self.video_path_column = video_path_column + self.video_folder = video_folder + self.sample_method = sample_method + self.num_sampled_frames = num_sampled_frames + self.num_sample_stride = num_sample_stride + + if video_path_list is not None: + self.video_path_list = video_path_list + self.metadata_df = pd.DataFrame({video_path_column: self.video_path_list}) + else: + self.video_path_list = get_video_path_list( + video_folder=video_folder, + video_metadata_path=video_metadata_path, + video_path_column=video_path_column + ) + + def __getitem__(self, index): + # video_path = os.path.join(self.video_folder, str(self.video_path_list[index])) + video_path = self.video_path_list[index] + try: + sample_args = (video_path, self.sample_method, self.num_sampled_frames, self.num_sample_stride) + sampled_frame_idx_list, sampled_frame_list = func_timeout( + VIDEO_READER_TIMEOUT, extract_frames, args=sample_args + ) + except FunctionTimedOut: + logger.warning(f"Read {video_path} timeout.") + return None + except Exception as e: + logger.warning(f"Failed to extract frames from video {video_path}. Error is {e}.") + return None + item = { + "video_path": Path(video_path).name, + "sampled_frame_idx": sampled_frame_idx_list, + "sampled_frame": sampled_frame_list, + } + + return item + + def __len__(self): + return len(self.video_path_list) + + +if __name__ == "__main__": + video_folder = "your_video_folder" + video_dataset = VideoDataset(video_folder=video_folder) + + video_dataloader = DataLoader( + video_dataset, batch_size=16, num_workers=16, collate_fn=collate_fn + ) + for idx, batch in enumerate(video_dataloader): + if len(batch) != 0: + print(batch["video_path"], batch["sampled_frame_idx"], len(batch["video_path"])) \ No newline at end of file diff --git a/easyanimate/video_caption/utils/video_utils.py b/easyanimate/video_caption/utils/video_utils.py new file mode 100644 index 0000000..eaf6200 --- /dev/null +++ b/easyanimate/video_caption/utils/video_utils.py @@ -0,0 +1,87 @@ +import gc +import os +import random +from contextlib import contextmanager +from pathlib import Path +from typing import List, Tuple, Optional + +import numpy as np +import pandas as pd +from decord import VideoReader +from PIL import Image + +ALL_VIDEO_EXT = set([".mp4", ".webm", ".mkv", ".avi", ".flv", ".mov"]) + + +def get_video_path_list( + video_folder: Optional[str]=None, + video_metadata_path: Optional[str]=None, + video_path_column: Optional[str]=None +) -> List[str]: + """Get all video (absolute) path list from the video folder or the video metadata file. + + Args: + video_folder (str): The absolute path of the folder (including sub-folders) containing all the required video files. + video_metadata_path (str): The absolute path of the video metadata file containing video path list. + video_path_column (str): The column/key for the corresponding video path in the video metadata file (csv/jsonl). + """ + if video_folder is None and video_metadata_path is None: + raise ValueError("Either the video_input or the video_metadata_path should be specified.") + if video_metadata_path is not None: + if video_metadata_path.endswith(".csv"): + if video_path_column is None: + raise ValueError("The video_path_column can not be None if provided a csv file.") + metadata_df = pd.read_csv(video_metadata_path) + video_path_list = metadata_df[video_path_column].tolist() + elif video_metadata_path.endswith(".jsonl"): + if video_path_column is None: + raise ValueError("The video_path_column can not be None if provided a jsonl file.") + metadata_df = pd.read_json(video_metadata_path, lines=True) + video_path_list = metadata_df[video_path_column].tolist() + elif video_metadata_path.endswith(".txt"): + with open(video_metadata_path, "r", encoding="utf-8") as f: + video_path_list = [line.strip() for line in f] + else: + raise ValueError("The video_metadata_path must end with `.csv`, `.jsonl` or `.txt`.") + if video_folder is not None: + video_path_list = [os.path.join(video_folder, video_path) for video_path in video_path_list] + return video_path_list + + if video_folder is not None: + video_path_list = [] + for ext in ALL_VIDEO_EXT: + video_path_list.extend(Path(video_folder).rglob(f"*{ext}")) + video_path_list = [str(video_path) for video_path in video_path_list] + return video_path_list + + +@contextmanager +def video_reader(*args, **kwargs): + """A context manager to solve the memory leak of decord. + """ + vr = VideoReader(*args, **kwargs) + try: + yield vr + finally: + del vr + gc.collect() + + +def extract_frames( + video_path: str, sample_method: str = "mid", num_sampled_frames: int = -1, sample_stride: int = -1 +) -> Optional[Tuple[List[int], List[Image.Image]]]: + with video_reader(video_path, num_threads=2) as vr: + if sample_method == "mid": + sampled_frame_idx_list = [len(vr) // 2] + elif sample_method == "uniform": + sampled_frame_idx_list = np.linspace(0, len(vr), num_sampled_frames, endpoint=False, dtype=int) + elif sample_method == "random": + clip_length = min(len(vr), (num_sampled_frames - 1) * sample_stride + 1) + start_idx = random.randint(0, len(vr) - clip_length) + sampled_frame_idx_list = np.linspace(start_idx, start_idx + clip_length - 1, num_sampled_frames, dtype=int) + else: + raise ValueError("The sample_method must be mid, uniform or random.") + sampled_frame_list = vr.get_batch(sampled_frame_idx_list).asnumpy() + sampled_frame_list = [Image.fromarray(frame) for frame in sampled_frame_list] + + return list(sampled_frame_idx_list), sampled_frame_list