add the dataset preprocessing pipeline

This commit is contained in:
hkunzhe
2024-05-27 22:26:03 +08:00
parent 234c157395
commit 6efc7de3f2
8 changed files with 811 additions and 3 deletions
+5 -1
View File
@@ -13,7 +13,11 @@ EasyAnimate uses multi-modal LLMs to generate captions for frames extracted from
cd EasyAnimate && pip install -r requirements.txt
# Install additional requirements for video caption.
cd easyanimate/video_caption && pip install -r requirements.txt
cd easyanimate/video_caption && pip install -r requirements.txt --extra-index-url https://huggingface.github.io/autogptq-index/whl/cu118/
# Use DDP instead of DP in EasyOCR detection.
site_pkg_path=$(python -c 'import site; print(site.getsitepackages()[0])')
cp -v easyocr_detection_patched.py $site_pkg_path/easyocr/detection.py
# We strongly recommend using Docker unless you can properly handle the dependency between vllm with torch(cuda).
```
@@ -5,6 +5,7 @@ import os
import pandas as pd
from accelerate import PartialState
from accelerate.utils import gather_object
from natsort import natsorted
from tqdm import tqdm
from torch.utils.data import DataLoader
@@ -86,6 +87,11 @@ def accelerate_inference(args, video_path_list):
elif args.image_caption_model_name == "Qwen-VL-Chat":
image_caption_model = QwenVLChat(device=device, quantized=args.image_caption_model_quantized)
# The workaround can be removed after https://github.com/huggingface/accelerate/pull/2781 is released.
index = len(video_path_list) - len(video_path_list) % state.num_processes
logger.info(f"Drop {len(video_path_list) % state.num_processes} videos to avoid duplicates in state.split_between_processes.")
video_path_list = video_path_list[:index]
if state.is_main_process:
os.makedirs(args.output_dir, exist_ok=True)
result_list = []
@@ -242,6 +248,8 @@ def main():
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))
# Sorting to guarantee the same result for each process.
video_path_list = natsorted(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:
@@ -0,0 +1,192 @@
import argparse
import gc
import os
from contextlib import contextmanager
from pathlib import Path
import cv2
import numpy as np
import pandas as pd
from joblib import Parallel, delayed
from natsort import natsorted
from tqdm import tqdm
from utils.logger import logger
from utils.video_utils import get_video_path_list
@contextmanager
def VideoCapture(video_path):
cap = cv2.VideoCapture(video_path)
try:
yield cap
finally:
cap.release()
del cap
gc.collect()
def compute_motion_score(video_path):
video_motion_scores = []
sampling_fps = 2
try:
with VideoCapture(video_path) as cap:
fps = cap.get(cv2.CAP_PROP_FPS)
valid_fps = min(max(sampling_fps, 1), fps)
frame_interval = int(fps / valid_fps)
total_frames = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
# if cannot get the second frame, use the last one
frame_interval = min(frame_interval, total_frames - 1)
prev_frame = None
frame_count = -1
while cap.isOpened():
ret, frame = cap.read()
frame_count += 1
if not ret:
break
# skip middle frames
if frame_count % frame_interval != 0:
continue
gray_frame = cv2.cvtColor(frame, cv2.COLOR_BGR2GRAY)
if prev_frame is None:
prev_frame = gray_frame
continue
flow = cv2.calcOpticalFlowFarneback(
prev_frame,
gray_frame,
None,
pyr_scale=0.5,
levels=3,
winsize=15,
iterations=3,
poly_n=5,
poly_sigma=1.2,
flags=0,
)
mag, _ = cv2.cartToPolar(flow[..., 0], flow[..., 1])
frame_motion_score = np.mean(mag)
video_motion_scores.append(frame_motion_score)
prev_frame = gray_frame
video_meta_info = {
"video_path": Path(video_path).name,
"motion_score": round(float(np.mean(video_motion_scores)), 5),
}
return video_meta_info
except Exception as e:
print(f"Compute motion score for video {video_path} with error: {e}.")
def parse_args():
parser = argparse.ArgumentParser(description="Compute the motion score of the videos.")
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)."
)
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("--saved_path", type=str, required=True, help="The save path to the output results (csv/jsonl).")
parser.add_argument("--saved_freq", type=int, default=100, help="The frequency to save the output results.")
parser.add_argument("--n_jobs", type=int, default=1, help="The number of concurrent processes.")
parser.add_argument(
"--asethetic_score_metadata_path", type=str, default=None, help="The path to the video quality metadata (csv/jsonl)."
)
parser.add_argument("--asethetic_score_threshold", type=float, default=4.0, help="The asethetic score threshold.")
parser.add_argument(
"--video_text_metadata_path", type=str, default=None, help="The path to the video text score metadata (csv/jsonl)."
)
parser.add_argument("--text_threshold", type=float, default=0.02, help="The text threshold.")
args = parser.parse_args()
return args
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, video_path) for video_path in saved_video_path_list]
video_path_list = list(set(video_path_list).difference(set(saved_video_path_list)))
# Sorting to guarantee the same result for each process.
video_path_list = natsorted(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.asethetic_score_metadata_path is not None:
if args.asethetic_score_metadata_path.endswith(".csv"):
asethetic_score_df = pd.read_csv(args.asethetic_score_metadata_path)
elif args.asethetic_score_metadata_path.endswith(".jsonl"):
asethetic_score_df = pd.read_json(args.asethetic_score_metadata_path, lines=True)
asethetic_score_df["aesthetic_score_mean"] = asethetic_score_df["aesthetic_score"].apply(lambda x: sum(x) / len(x))
filtered_asethetic_score_df = asethetic_score_df[asethetic_score_df["aesthetic_score_mean"] < args.asethetic_score_threshold]
filtered_video_path_list = filtered_asethetic_score_df[args.video_path_column].tolist()
filtered_video_path_list = [os.path.join(args.video_folder, video_path) for video_path in filtered_video_path_list]
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
# Sorting to guarantee the same result for each process.
video_path_list = natsorted(video_path_list)
logger.info(f"Load {args.asethetic_score_metadata_path} and filter {len(filtered_video_path_list)} videos.")
if args.text_score_metadata_path is not None:
if args.text_score_metadata_path.endswith(".csv"):
text_score_df = pd.read_csv(args.text_score_metadata_path)
elif args.text_score_metadata_path.endswith(".jsonl"):
text_score_df = pd.read_json(args.text_score_metadata_path, lines=True)
text_score_df["aesthetic_score_mean"] = text_score_df["aesthetic_score"].apply(lambda x: sum(x) / len(x))
filtered_text_score_df = text_score_df[text_score_df["text_score"] > args.text_score_threshold]
filtered_video_path_list = filtered_text_score_df[args.video_path_column].tolist()
filtered_video_path_list = [os.path.join(args.video_folder, video_path) for video_path in filtered_video_path_list]
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
# Sorting to guarantee the same result for each process.
video_path_list = natsorted(video_path_list)
logger.info(f"Load {args.text_score_metadata_path} and filter {len(filtered_video_path_list)} videos.")
for i in tqdm(range(0, len(video_path_list), args.saved_freq)):
result_list = Parallel(n_jobs=args.n_jobs, backend="threading")(
delayed(compute_motion_score)(video_path) for video_path in tqdm(video_path_list[i: i + args.saved_freq])
)
result_list = [result for result in result_list if result is not None]
if len(result_list) == 0:
continue
result_df = pd.DataFrame(result_list)
if args.saved_path.endswith(".csv"):
header = False if os.path.exists(args.saved_path) else True
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}.")
if __name__ == "__main__":
main()
@@ -0,0 +1,174 @@
import argparse
import os
from pathlib import Path
import easyocr
import numpy as np
import pandas as pd
from accelerate import PartialState
from accelerate.utils import gather_object
from natsort import natsorted
from tqdm import tqdm
from utils.logger import logger
from utils.video_utils import extract_frames, get_video_path_list
# @contextmanager
# def video_reader(*args, **kwargs):
# vr = VideoReader(*args, **kwargs)
# try:
# yield vr
# finally:
# del vr
# gc.collect()
# def extract_mid_frame(video_path: str):
# with video_reader(video_path, num_threads=2) as vr:
# middle_frame_index = len(vr) // 2
# middle_frame = vr[middle_frame_index].asnumpy()
# return [middle_frame_index], [middle_frame]
def triangle_area(p1, p2, p3):
"""Compute the triangle area according to its coordinates.
"""
x1, y1 = p1
x2, y2 = p2
x3, y3 = p3
tri_area = 0.5 * np.abs(x1 * y2 + x2 * y3 + x3 * y1 - x2 * y1 - x3 * y2 - x1 * y3)
return tri_area
def compute_text_score(video_path, ocr_reader):
_, images = extract_frames(video_path, sample_method="mid")
frame_ocr_area_ratios = []
for image in images:
# horizontal detected results and free-form detected
horizontal_list, free_list = ocr_reader.detect(np.asarray(image))
width, height = image.shape[0], image.shape[1]
total_area = width * height
# rectangles
rect_area = 0
for xmin, xmax, ymin, ymax in horizontal_list[0]:
if xmax < xmin or ymax < ymin:
continue
rect_area += (xmax - xmin) * (ymax - ymin)
# free-form
quad_area = 0
try:
for points in free_list[0]:
triangle1 = points[:3]
quad_area += triangle_area(*triangle1)
triangle2 = points[3:] + [points[0]]
quad_area += triangle_area(*triangle2)
except:
quad_area = 0
text_area = rect_area + quad_area
frame_ocr_area_ratios.append(text_area / total_area)
video_meta_info = {
"video_path": Path(video_path).name,
"text_score": round(np.mean(frame_ocr_area_ratios), 5),
}
return video_meta_info
def parse_args():
parser = argparse.ArgumentParser(description="Compute the text score of the middle frame in the videos.")
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)."
)
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("--saved_path", type=str, required=True, help="The save path to the output results (csv/jsonl).")
parser.add_argument("--saved_freq", type=int, default=100, help="The frequency to save the output results.")
args = parser.parse_args()
return args
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, video_path) for video_path in saved_video_path_list]
video_path_list = list(set(video_path_list).difference(set(saved_video_path_list)))
# Sorting to guarantee the same result for each process.
video_path_list = natsorted(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.")
state = PartialState()
ocr_reader = easyocr.Reader(
lang_list=["en", "ch_sim"],
gpu=state.device,
recognizer=False,
verbose=False,
model_storage_directory="/mnt/nas/huangkunzhe.hkz/code/video-caption/models/",
# https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/video_caption/easyocr/craft_mlt_25k.pth
)
# The workaround can be removed after https://github.com/huggingface/accelerate/pull/2781 is released.
index = len(video_path_list) - len(video_path_list) % state.num_processes
logger.info(f"Drop {len(video_path_list) % state.num_processes} videos to avoid duplicates in state.split_between_processes.")
video_path_list = video_path_list[:index]
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)):
video_meta_info = compute_text_score(video_path, ocr_reader)
result_list.append(video_meta_info)
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"):
header = False if os.path.exists(args.saved_path) else True
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_list = []
state.wait_for_everyone()
gathered_result_list = gather_object(result_list)
if state.is_main_process:
logger.info(len(gathered_result_list))
if len(gathered_result_list) != 0:
result_df = pd.DataFrame(gathered_result_list)
if args.saved_path.endswith(".csv"):
header = False if os.path.exists(args.saved_path) else True
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,114 @@
"""Modified from https://github.com/JaidedAI/EasyOCR/blob/803b907/easyocr/detection.py.
1. Disable DataParallel.
"""
import torch
import torch.backends.cudnn as cudnn
from torch.autograd import Variable
from PIL import Image
from collections import OrderedDict
import cv2
import numpy as np
from .craft_utils import getDetBoxes, adjustResultCoordinates
from .imgproc import resize_aspect_ratio, normalizeMeanVariance
from .craft import CRAFT
def copyStateDict(state_dict):
if list(state_dict.keys())[0].startswith("module"):
start_idx = 1
else:
start_idx = 0
new_state_dict = OrderedDict()
for k, v in state_dict.items():
name = ".".join(k.split(".")[start_idx:])
new_state_dict[name] = v
return new_state_dict
def test_net(canvas_size, mag_ratio, net, image, text_threshold, link_threshold, low_text, poly, device, estimate_num_chars=False):
if isinstance(image, np.ndarray) and len(image.shape) == 4: # image is batch of np arrays
image_arrs = image
else: # image is single numpy array
image_arrs = [image]
img_resized_list = []
# resize
for img in image_arrs:
img_resized, target_ratio, size_heatmap = resize_aspect_ratio(img, canvas_size,
interpolation=cv2.INTER_LINEAR,
mag_ratio=mag_ratio)
img_resized_list.append(img_resized)
ratio_h = ratio_w = 1 / target_ratio
# preprocessing
x = [np.transpose(normalizeMeanVariance(n_img), (2, 0, 1))
for n_img in img_resized_list]
x = torch.from_numpy(np.array(x))
x = x.to(device)
# forward pass
with torch.no_grad():
y, feature = net(x)
boxes_list, polys_list = [], []
for out in y:
# make score and link map
score_text = out[:, :, 0].cpu().data.numpy()
score_link = out[:, :, 1].cpu().data.numpy()
# Post-processing
boxes, polys, mapper = getDetBoxes(
score_text, score_link, text_threshold, link_threshold, low_text, poly, estimate_num_chars)
# coordinate adjustment
boxes = adjustResultCoordinates(boxes, ratio_w, ratio_h)
polys = adjustResultCoordinates(polys, ratio_w, ratio_h)
if estimate_num_chars:
boxes = list(boxes)
polys = list(polys)
for k in range(len(polys)):
if estimate_num_chars:
boxes[k] = (boxes[k], mapper[k])
if polys[k] is None:
polys[k] = boxes[k]
boxes_list.append(boxes)
polys_list.append(polys)
return boxes_list, polys_list
def get_detector(trained_model, device='cpu', quantize=True, cudnn_benchmark=False):
net = CRAFT()
if device == 'cpu':
net.load_state_dict(copyStateDict(torch.load(trained_model, map_location=device)))
if quantize:
try:
torch.quantization.quantize_dynamic(net, dtype=torch.qint8, inplace=True)
except:
pass
else:
net.load_state_dict(copyStateDict(torch.load(trained_model, map_location=device)))
# net = torch.nn.DataParallel(net).to(device)
net = net.to(device)
cudnn.benchmark = cudnn_benchmark
net.eval()
return net
def get_textbox(detector, image, canvas_size, mag_ratio, text_threshold, link_threshold, low_text, poly, device, optimal_num_chars=None, **kwargs):
result = []
estimate_num_chars = optimal_num_chars is not None
bboxes_list, polys_list = test_net(canvas_size, mag_ratio, detector,
image, text_threshold,
link_threshold, low_text, poly,
device, estimate_num_chars)
if estimate_num_chars:
polys_list = [[p for p, _ in sorted(polys, key=lambda x: abs(optimal_num_chars - x[1]))]
for polys in polys_list]
for polys in polys_list:
single_img_result = []
for i, box in enumerate(polys):
poly = np.array(box).astype(np.int32).reshape((-1))
single_img_result.append(poly)
result.append(single_img_result)
return result
+5 -2
View File
@@ -1,6 +1,9 @@
--extra-index-url https://huggingface.github.io/autogptq-index/whl/cu118/
auto_gptq==0.6.0
pandas>=2.0.0
vllm==0.3.3
sglang[srt]==0.1.13
func_timeout
func_timeout
easyocr==1.7.1
git+https://github.com/openai/CLIP.git
natsort
joblib
@@ -0,0 +1,130 @@
import os
from typing import List
import clip
import torch
import torch.nn as nn
import torch.nn.functional as F
from PIL import Image
from torchvision.datasets.utils import download_url
from transformers import AutoModel, AutoProcessor
# All metrics.
__all__ = ["AestheticScore", "CLIPScore"]
_MODELS = {
"CLIP_ViT-L/14": "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/video_caption/clip/ViT-L-14.pt",
"Aesthetics_V2": "https://pai-aigc-photog.oss-cn-hangzhou.aliyuncs.com/easyanimate/video_caption/clip/sac%2Blogos%2Bava1-l14-linearMSE.pth",
}
_MD5 = {
"CLIP_ViT-L/14": "096db1af569b284eb76b3881534822d9",
"Aesthetics_V2": "b1047fd767a00134b8fd6529bf19521a",
}
# if you changed the MLP architecture during training, change it also here:
class _MLP(nn.Module):
def __init__(self, input_size):
super().__init__()
self.input_size = input_size
self.layers = nn.Sequential(
nn.Linear(self.input_size, 1024),
# nn.ReLU(),
nn.Dropout(0.2),
nn.Linear(1024, 128),
# nn.ReLU(),
nn.Dropout(0.2),
nn.Linear(128, 64),
# nn.ReLU(),
nn.Dropout(0.1),
nn.Linear(64, 16),
# nn.ReLU(),
nn.Linear(16, 1),
)
def forward(self, x):
return self.layers(x)
class AestheticScore:
"""Compute LAION Aesthetics Score V2 based on openai/clip. Note that the default
inference dtype with GPUs is fp16 in openai/clip.
Ref:
1. https://github.com/christophschuhmann/improved-aesthetic-predictor/blob/main/simple_inference.py.
2. https://github.com/openai/CLIP/issues/30.
"""
def __init__(self, root: str = "~/.cache/clip", device: str = "cpu"):
# The CLIP model is loaded in the evaluation mode.
self.root = os.path.expanduser(root)
if not os.path.exists(self.root):
os.makedirs(self.root)
filename = "ViT-L-14.pt"
download_url(_MODELS["CLIP_ViT-L/14"], self.root, filename=filename, md5=_MD5["CLIP_ViT-L/14"])
self.clip_model, self.preprocess = clip.load(os.path.join(self.root, filename), device=device)
self.device = device
self._load_mlp()
def _load_mlp(self):
filename = "sac+logos+ava1-l14-linearMSE.pth"
download_url(_MODELS["Aesthetics_V2"], self.root, filename=filename, md5=_MD5["Aesthetics_V2"])
state_dict = torch.load(os.path.join(self.root, filename))
self.mlp = _MLP(768)
self.mlp.load_state_dict(state_dict)
self.mlp.to(self.device)
self.mlp.eval()
def __call__(self, images: List[Image.Image], texts=None) -> List[float]:
with torch.no_grad():
images = torch.stack([self.preprocess(image) for image in images]).to(self.device)
image_embs = F.normalize(self.clip_model.encode_image(images))
scores = self.mlp(image_embs.float()) # torch.float16 -> torch.float32, [N, 1]
return scores.squeeze().tolist()
def __repr__(self) -> str:
return "aesthetic_score"
class CLIPScore:
"""Compute CLIP scores for image-text pairs based on huggingface/transformers."""
def __init__(
self,
model_name_or_path: str = "openai/clip-vit-large-patch14",
torch_dtype=torch.float16,
device: str = "cpu",
):
self.model = AutoModel.from_pretrained(model_name_or_path, torch_dtype=torch_dtype).eval().to(device)
self.processor = AutoProcessor.from_pretrained(model_name_or_path)
self.torch_dtype = torch_dtype
self.device = device
def __call__(self, images: List[Image.Image], texts: List[str]) -> List[float]:
assert len(images) == len(texts)
image_inputs = self.processor(images=images, return_tensors="pt") # {"pixel_values": }
if self.torch_dtype == torch.float16:
image_inputs["pixel_values"] = image_inputs["pixel_values"].half()
text_inputs = self.processor(text=texts, return_tensors="pt", padding=True, truncation=True) # {"inputs_id": }
image_inputs, text_inputs = image_inputs.to(self.device), text_inputs.to(self.device)
with torch.no_grad():
image_embs = F.normalize(self.model.get_image_features(**image_inputs))
text_embs = F.normalize(self.model.get_text_features(**text_inputs))
scores = text_embs @ image_embs.T # [N, N]
return scores.diagonal().tolist()
def __repr__(self) -> str:
return "clip_score"
if __name__ == "__main__":
aesthetic_score = AestheticScore(device="cuda")
clip_score = CLIPScore(device="cuda")
paths = ["demo/splash_cl2_midframe.jpg"] * 3
texts = ["a joker", "a woman", "a man"]
images = [Image.open(p).convert("RGB") for p in paths]
print(aesthetic_score(images))
print(clip_score(images, texts))
@@ -0,0 +1,183 @@
import argparse
import re
import os
import pandas as pd
from accelerate import PartialState
from accelerate.utils import gather_object
from natsort import natsorted
from tqdm import tqdm
from torch.utils.data import DataLoader
import utils.image_evaluator as image_evaluator
from utils.logger import logger
from utils.video_dataset import VideoDataset, collate_fn
from utils.video_utils import get_video_path_list
def camel2snake(s: str) -> str:
"""Convert camel case to snake case."""
if not re.match("^[A-Z]+$", s):
pattern = re.compile(r"(?<!^)(?=[A-Z])")
return pattern.sub("_", s).lower()
return s
def parse_args():
parser = argparse.ArgumentParser(description="Compute scores of uniform sampled frames from videos.")
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)."
)
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=None,
help="The column contains the caption.",
)
parser.add_argument(
"--num_sampled_frames",
type=int,
default=4,
help="num_sampled_frames",
)
parser.add_argument("--metrics", nargs="+", type=str, required=True, help="The evaluation metric(s) for generated images.")
parser.add_argument(
"--batch_size",
type=int,
default=10,
required=False,
help="The batch size for the video dataset.",
)
parser.add_argument(
"--output_dir",
type=str,
required=True,
help="The directory to creat 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.")
parser.add_argument("--resume", default=False, action="store_true", help="Whether to resume from the saved_path.")
args = parser.parse_args()
return args
def main():
args = parse_args()
assert args.batch_size > 1
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.")
caption_list = None
if args.video_metadata_path is not None and args.caption_column is not None:
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.")
caption_list = video_metadata_df[args.caption_column].tolist()
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, video_path) for video_path in saved_video_path_list]
video_path_list = list(set(video_path_list).difference(set(saved_video_path_list)))
# Sorting to guarantee the same result for each process.
video_path_list = natsorted(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.")
logger.info("Initializing evaluator metrics...")
state = PartialState()
metric_fns = [getattr(image_evaluator, metric)(device=state.device) for metric in args.metrics]
# The workaround can be removed after https://github.com/huggingface/accelerate/pull/2781 is released.
index = len(video_path_list) - len(video_path_list) % state.num_processes
logger.info(f"Drop {len(video_path_list) % state.num_processes} videos to avoid duplicates in state.split_between_processes.")
video_path_list = video_path_list[:index]
result_dict = {args.video_path_column: [], "sample_frame_idx": []}
for metric in args.metrics:
result_dict[camel2snake(metric)] = []
with state.split_between_processes(video_path_list) as splitted_video_path_list:
video_dataset = VideoDataset(
video_path_list=splitted_video_path_list,
sample_method="extract_uniform_frames",
num_sampled_frames=args.num_sampled_frames
)
video_loader = DataLoader(video_dataset, batch_size=args.batch_size, num_workers=4, collate_fn=collate_fn)
for idx, batch in enumerate(tqdm(video_loader)):
if len(batch) == 0:
continue
batch_video_path = batch[args.video_path_column]
result_dict["sample_frame_idx"].extend(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])
batch_caption = None
if caption_list is not None:
batch_caption = caption_list[i : i + args.batch_size]
# Compute the frame quality.
for i, metric in enumerate(args.metrics):
# [batch_size * num_sampled_frames] => [batch_size, num_sampled_frames]
quality_scores = metric_fns[i](batch_frame, batch_caption)
quality_scores = [round(score, 5) for score in quality_scores]
quality_scores = [quality_scores[j:j + args.num_sampled_frames] for j in range(0, len(quality_scores), args.num_sampled_frames)]
result_dict[camel2snake(metric)].extend(quality_scores)
saved_video_path_list = [os.path.basename(video_path) for video_path in batch_video_path]
result_dict[args.video_path_column].extend(saved_video_path_list)
# Save the metadata in the main process every saved_freq.
if (idx != 0) and (idx % args.saved_freq == 0):
state.wait_for_everyone()
gathered_result_dict = {k: gather_object(v) for k, v in result_dict.items()}
if state.is_main_process:
result_df = pd.DataFrame(gathered_result_dict)
if args.saved_path.endswith(".csv"):
header = False if os.path.exists(args.saved_path) else True
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}.")
for k in result_dict.keys():
result_dict[k] = []
# Wait for all processes to finish and gather the final result.
state.wait_for_everyone()
gathered_result_dict = {k: gather_object(v) for k, v in result_dict.items()}
# Save the metadata in the main process.
if state.is_main_process:
result_df = pd.DataFrame(gathered_result_dict)
if len(gathered_result_dict[args.video_path_column]) != 0:
result_df = pd.DataFrame(gathered_result_dict)
if args.saved_path.endswith(".csv"):
header = False if os.path.exists(args.saved_path) else True
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()