add video caption

This commit is contained in:
huangkunzhe.hkz
2024-04-15 16:25:11 +08:00
parent d6ec701713
commit a4008a6127
8 changed files with 828 additions and 0 deletions
+59
View File
@@ -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.
@@ -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()
@@ -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()
@@ -0,0 +1,5 @@
pandas>=2.0.0
auto_gptq
vllm
sglang[srt]
func_timeout
@@ -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("<ImageHere>"):
prompt = "<ImageHere>" + 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)
+36
View File
@@ -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")
@@ -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"]))
@@ -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