Fix train reward lora and video caption (#165)
This commit is contained in:
@@ -284,7 +284,6 @@ class Cross_model(nn.Module):
|
||||
context_tokens,
|
||||
mask
|
||||
):
|
||||
print(mask.dtype)
|
||||
for cross_attn, self_attn_ff in self.layers:
|
||||
query_tokens = cross_attn(query_tokens, context_tokens,mask)
|
||||
query_tokens = self_attn_ff(query_tokens)
|
||||
|
||||
@@ -144,6 +144,7 @@ We support batched inference with local LLMs or OpenAI compatible server based o
|
||||
--model_name /path/to/your_llm \
|
||||
--prompt prompt/beautiful_prompt.txt \
|
||||
--prefix '"detailed description": ' \
|
||||
--max_retry_count 10 \
|
||||
--saved_path datasets/beautiful_prompt.jsonl \
|
||||
--saved_freq 1
|
||||
```
|
||||
|
||||
@@ -129,6 +129,7 @@ CAPTION_MODEL_PATH=/PATH/TO/INTERNVL2_MODEL REWRITE_MODEL_PATH=/PATH/TO/REWRITE_
|
||||
--model_name /path/to/your_llm \
|
||||
--prompt prompt/beautiful_prompt.txt \
|
||||
--prefix '"detailed description": ' \
|
||||
--max_retry_count 10 \
|
||||
--saved_path datasets/beautiful_prompt.jsonl \
|
||||
--saved_freq 1
|
||||
```
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import argparse
|
||||
import os
|
||||
import re
|
||||
from copy import deepcopy
|
||||
|
||||
import pandas as pd
|
||||
import torch
|
||||
@@ -74,6 +75,18 @@ def parse_args():
|
||||
required=True,
|
||||
help="The prefix to extract the output from LLMs.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--answer_template",
|
||||
type=str,
|
||||
default="",
|
||||
help="The anwer template in the prompt. If specified, rewritten results same as the answer template will be removed.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--max_retry_count",
|
||||
type=int,
|
||||
default=1,
|
||||
help="The maximum retry count to ensure outputs with the valid format from LLMs.",
|
||||
)
|
||||
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=1, help="The frequency to save the output results.")
|
||||
|
||||
@@ -102,14 +115,28 @@ def main():
|
||||
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)
|
||||
|
||||
# Remove previous rewritten results same as the answer template.
|
||||
if args.answer_template != "":
|
||||
prev_nums = len(saved_metadata_df)
|
||||
saved_metadata_df = saved_metadata_df[
|
||||
~saved_metadata_df[args.caption_column].str.contains(args.answer_template, case=False, na=False)
|
||||
]
|
||||
logger.info(
|
||||
f"Remove {prev_nums - len(saved_metadata_df)} rewritten results same as the answer template "
|
||||
f"from {args.saved_path}."
|
||||
)
|
||||
if args.saved_path.endswith(".csv"):
|
||||
saved_metadata_df.to_csv(args.saved_path, index=False)
|
||||
elif args.saved_path.endswith(".jsonl"):
|
||||
saved_metadata_df.to_json(args.saved_path, orient="records", lines=True)
|
||||
|
||||
# Filter out the unprocessed video-caption pairs by setting the indicator=True.
|
||||
merged_df = video_metadata_df.merge(saved_metadata_df, on=args.video_path_column, how="outer", indicator=True)
|
||||
video_metadata_df = merged_df[merged_df["_merge"] == "left_only"]
|
||||
# Sorting to guarantee the same result for each process.
|
||||
video_metadata_df = video_metadata_df.iloc[index_natsorted(video_metadata_df[args.video_path_column])].reset_index(
|
||||
drop=True
|
||||
)
|
||||
video_metadata_df = video_metadata_df.iloc[index_natsorted(video_metadata_df[args.video_path_column])]
|
||||
video_metadata_df = video_metadata_df.reset_index(drop=True)
|
||||
logger.info(
|
||||
f"Resume from {args.saved_path}: {len(saved_metadata_df)} processed and {len(video_metadata_df)} to be processed."
|
||||
)
|
||||
@@ -119,6 +146,9 @@ def main():
|
||||
args.prompt = "".join(f.readlines())
|
||||
logger.info(f"Prompt: {args.prompt}")
|
||||
|
||||
if args.max_retry_count < 1:
|
||||
raise ValueError(f"The max_retry_count {args.max_retry_count} must be greater than 0.")
|
||||
|
||||
if args.video_path_column is not None:
|
||||
video_path_list = video_metadata_df[args.video_path_column].tolist()
|
||||
if args.caption_column in video_metadata_df.columns:
|
||||
@@ -162,27 +192,38 @@ def main():
|
||||
]
|
||||
text = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True)
|
||||
batch_prompt.append(text)
|
||||
|
||||
cur_retry_count = 0
|
||||
while cur_retry_count < args.max_retry_count:
|
||||
if len(batch_prompt) == 0:
|
||||
break
|
||||
|
||||
batch_output = llm.generate(batch_prompt, sampling_params)
|
||||
batch_output = [output.outputs[0].text.rstrip() for output in batch_output]
|
||||
batch_output = [extract_output(output, prefix=args.prefix) for output in batch_output]
|
||||
batch_result = []
|
||||
batch_output = llm.generate(batch_prompt, sampling_params)
|
||||
batch_output = [output.outputs[0].text.rstrip() for output in batch_output]
|
||||
if args.prefix is not None:
|
||||
batch_output = [extract_output(output, args.prefix) for output in batch_output]
|
||||
|
||||
# Filter out data that does not meet the output format.
|
||||
batch_result = []
|
||||
if args.video_path_column is not None:
|
||||
for video_path, output in zip(batch_video_path, batch_output):
|
||||
# Filter out data that does not meet the output format to retry.
|
||||
retry_batch_video_path, retry_batch_prompt = [], []
|
||||
for (video_path, prompt, output) in zip(batch_video_path, batch_prompt, batch_output):
|
||||
if output is not None:
|
||||
batch_result.append((video_path, output))
|
||||
batch_video_path, batch_output = zip(*batch_result)
|
||||
|
||||
result_dict[args.video_path_column].extend(batch_video_path)
|
||||
result_dict[args.caption_column].extend(batch_output)
|
||||
else:
|
||||
for output in batch_output:
|
||||
if output is not None:
|
||||
batch_result.append(output)
|
||||
|
||||
result_dict[args.caption_column].extend(batch_result)
|
||||
else:
|
||||
retry_batch_video_path.append(video_path)
|
||||
retry_batch_prompt.append(prompt)
|
||||
|
||||
if len(batch_result) != 0:
|
||||
batch_video_path, batch_output = zip(*batch_result)
|
||||
result_dict[args.video_path_column].extend(deepcopy(batch_video_path))
|
||||
result_dict[args.caption_column].extend(deepcopy(batch_output))
|
||||
|
||||
batch_video_path, batch_prompt = retry_batch_video_path, retry_batch_prompt
|
||||
cur_retry_count += 1
|
||||
logger.info(
|
||||
f"Current retry count/Maximum retry count: {cur_retry_count}/{args.max_retry_count}.: "
|
||||
f"Retrying {len(batch_prompt)} prompts with invalid output format."
|
||||
)
|
||||
|
||||
# Save the metadata every args.saved_freq.
|
||||
if (i // args.batch_size) % args.saved_freq == 0 or (i + 1) * args.batch_size >= len(sampled_frame_caption_list):
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
import argparse
|
||||
import ast
|
||||
import gc
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from pathlib import Path
|
||||
|
||||
import cv2
|
||||
import numpy as np
|
||||
@@ -76,11 +74,11 @@ def compute_motion_score(video_path):
|
||||
video_motion_scores.append(frame_motion_score)
|
||||
prev_frame = gray_frame
|
||||
|
||||
video_meta_info = {
|
||||
"video_path": Path(video_path).name,
|
||||
motion_score_result = {
|
||||
"video_path": video_path,
|
||||
"motion_score": round(float(np.mean(video_motion_scores)), 5),
|
||||
}
|
||||
return video_meta_info
|
||||
return motion_score_result
|
||||
|
||||
except Exception as e:
|
||||
print(f"Compute motion score for video {video_path} with error: {e}.")
|
||||
@@ -168,17 +166,27 @@ def main():
|
||||
min_aesthetic_score_siglip=args.min_aesthetic_score_siglip,
|
||||
text_score_metadata_path=args.text_score_metadata_path,
|
||||
min_text_score=args.min_text_score,
|
||||
semantic_consistency_score_metadata_path=args.semantic_consistency_score_metadata_path,
|
||||
min_semantic_consistency_score=args.min_semantic_consistency_score,
|
||||
video_path_column=args.video_path_column
|
||||
)
|
||||
video_path_list = [os.path.join(args.video_folder, video_path) for video_path in video_path_list]
|
||||
# Sorting to guarantee the same result for each process.
|
||||
video_path_list = natsorted(video_path_list)
|
||||
logger.info(f"{len(video_path_list)} videos are to be processed.")
|
||||
|
||||
for i in tqdm(range(0, len(video_path_list), args.saved_freq)):
|
||||
result_list = Parallel(n_jobs=args.n_jobs)(
|
||||
# Get motion score result for each video asynchronously.
|
||||
motion_score_result_list = Parallel(n_jobs=args.n_jobs)(
|
||||
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]
|
||||
result_list = []
|
||||
for motion_score_result in motion_score_result_list:
|
||||
if motion_score_result is not None:
|
||||
video_path = motion_score_result["video_path"]
|
||||
if args.video_folder != "":
|
||||
video_path = os.path.relpath(video_path, args.video_folder)
|
||||
result_list.append({args.video_path_column: video_path, "motion_score": motion_score_result["motion_score"]})
|
||||
if len(result_list) == 0:
|
||||
continue
|
||||
|
||||
|
||||
@@ -130,7 +130,6 @@ def main():
|
||||
max_motion_score=args.max_motion_score,
|
||||
video_path_column=args.video_path_column
|
||||
)
|
||||
video_path_list = [os.path.join(args.video_folder, video_path) for video_path in video_path_list]
|
||||
# Sorting to guarantee the same result for each process.
|
||||
video_path_list = natsorted(video_path_list)
|
||||
|
||||
@@ -201,7 +200,7 @@ def main():
|
||||
result_dict["sample_frame_idx"].extend(batch["sampled_frame_idx"])
|
||||
|
||||
# Save the metadata in the main process every saved_freq.
|
||||
if (idx != 0) and (idx % args.saved_freq == 0 or idx == len(video_loader) - 1):
|
||||
if (idx % args.saved_freq) == 0 or idx == len(video_loader) - 1:
|
||||
state.wait_for_everyone()
|
||||
gathered_result_dict = {k: gather_object(v) for k, v in result_dict.items()}
|
||||
if state.is_main_process and len(gathered_result_dict[args.video_path_column]) != 0:
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
import argparse
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import easyocr
|
||||
import numpy as np
|
||||
@@ -47,8 +46,8 @@ def triangle_area(p1, p2, p3):
|
||||
return tri_area
|
||||
|
||||
|
||||
def compute_text_score(video_path, ocr_reader):
|
||||
_, images = extract_frames(video_path, sample_method="mid")
|
||||
def compute_text_score(video_path, ocr_reader, sample_method="mid", num_sampled_frames=1):
|
||||
_, images = extract_frames(video_path, sample_method=sample_method, num_sampled_frames=num_sampled_frames)
|
||||
images = [np.array(image) for image in images]
|
||||
|
||||
frame_ocr_area_ratios = []
|
||||
@@ -78,12 +77,9 @@ def compute_text_score(video_path, ocr_reader):
|
||||
|
||||
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),
|
||||
}
|
||||
text_score = round(np.mean(frame_ocr_area_ratios), 5)
|
||||
|
||||
return video_meta_info
|
||||
return text_score
|
||||
|
||||
|
||||
def parse_args():
|
||||
@@ -98,6 +94,17 @@ def parse_args():
|
||||
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(
|
||||
"--frame_sample_method",
|
||||
type=str,
|
||||
default="mid",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num_sampled_frames",
|
||||
type=int,
|
||||
default=1,
|
||||
help="num_sampled_frames",
|
||||
)
|
||||
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=1, help="The frequency to save the output results.")
|
||||
|
||||
@@ -197,7 +204,18 @@ def main():
|
||||
with state.split_between_processes(video_path_list) as splitted_video_path_list:
|
||||
for i, video_path in enumerate(tqdm(splitted_video_path_list)):
|
||||
try:
|
||||
video_meta_info = compute_text_score(video_path, ocr_reader)
|
||||
text_score = compute_text_score(
|
||||
video_path,
|
||||
ocr_reader,
|
||||
sample_method=args.frame_sample_method,
|
||||
num_sampled_frames=args.num_sampled_frames,
|
||||
)
|
||||
video_meta_info = {}
|
||||
if args.video_folder == "":
|
||||
video_meta_info[args.video_path_column] = video_path
|
||||
else:
|
||||
video_meta_info[args.video_path_column] = os.path.relpath(video_path, args.video_folder)
|
||||
video_meta_info["text_score"] = text_score
|
||||
result_list.append(video_meta_info)
|
||||
except Exception as e:
|
||||
logger.warning(f"Compute text score for video {video_path} with error: {e}.")
|
||||
|
||||
@@ -15,14 +15,14 @@ def cutscene_detection_star(args):
|
||||
return cutscene_detection(*args)
|
||||
|
||||
|
||||
def cutscene_detection(video_path, saved_path, cutscene_threshold=27, min_scene_len=15):
|
||||
def cutscene_detection(video_path, video_folder, saved_path, cutscene_threshold=27, min_scene_len=15):
|
||||
try:
|
||||
if os.path.exists(saved_path):
|
||||
logger.info(f"{video_path} has been processed.")
|
||||
return
|
||||
# Use PyAV as the backend to avoid (to some exent) containing the last frame of the previous scene.
|
||||
# https://github.com/Breakthrough/PySceneDetect/issues/279#issuecomment-2152596761.
|
||||
video = open_video(video_path, backend="pyav")
|
||||
video = open_video(os.path.join(video_folder, video_path), backend="pyav")
|
||||
frame_rate, frame_size = video.frame_rate, video.frame_size
|
||||
duration = deepcopy(video.duration)
|
||||
|
||||
@@ -50,12 +50,14 @@ def cutscene_detection(video_path, saved_path, cutscene_threshold=27, min_scene_
|
||||
|
||||
timecode_list = [(frame_timecode_tuple[0].get_timecode(), frame_timecode_tuple[1].get_timecode()) for frame_timecode_tuple in output_scene_list]
|
||||
meta_scene = [{
|
||||
"video_path": Path(video_path).name,
|
||||
"video_path": video_path,
|
||||
"timecode_list": timecode_list,
|
||||
"fram_rate": frame_rate,
|
||||
"frame_size": frame_size,
|
||||
"duration": str(duration) # __repr__
|
||||
}]
|
||||
if not os.path.exists(Path(saved_path).parent):
|
||||
os.makedirs(Path(saved_path).parent, exist_ok=True)
|
||||
pd.DataFrame(meta_scene).to_json(saved_path, orient="records", lines=True)
|
||||
except Exception as e:
|
||||
logger.warning(f"Cutscene detection with {video_path} failed. Error is: {e}.")
|
||||
@@ -74,20 +76,18 @@ if __name__ == "__main__":
|
||||
)
|
||||
parser.add_argument("--video_folder", type=str, default="", help="The video folder.")
|
||||
parser.add_argument("--saved_folder", type=str, required=True, help="The save path to the output results (csv/jsonl).")
|
||||
parser.add_argument("--cutscene_threshold", type=int, default=27, help="The threshold of ContentDetector.")
|
||||
parser.add_argument("--n_jobs", type=int, default=1, help="The number of processes.")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
metadata_df = pd.read_json(args.video_metadata_path, lines=True)
|
||||
video_path_list = metadata_df[args.video_path_column].tolist()
|
||||
video_path_list = [os.path.join(args.video_folder, video_path) for video_path in video_path_list]
|
||||
|
||||
if not os.path.exists(args.saved_folder):
|
||||
os.makedirs(args.saved_folder, exist_ok=True)
|
||||
# The glob can be slow when there are many small jsonl files.
|
||||
saved_path_list = [os.path.join(args.saved_folder, Path(video_path).stem + ".jsonl") for video_path in video_path_list]
|
||||
saved_path_list = [os.path.join(args.saved_folder, Path(video_path).with_suffix(".jsonl")) for video_path in video_path_list]
|
||||
args_list = [
|
||||
(video_path, saved_path)
|
||||
(video_path, args.video_folder, saved_path, args.cutscene_threshold)
|
||||
for video_path, saved_path in zip(video_path_list, saved_path_list)
|
||||
]
|
||||
# Since the length of the video is not uniform, the gather operation is not performed.
|
||||
|
||||
@@ -143,7 +143,6 @@ def main():
|
||||
else:
|
||||
raise ValueError("The video_metadata_path must end with .csv or .jsonl.")
|
||||
video_path_list = video_metadata_df[args.video_path_column].tolist()
|
||||
video_path_list = [os.path.basename(video_path) for video_path in video_path_list]
|
||||
|
||||
if not (args.saved_path.endswith(".csv") or args.saved_path.endswith(".jsonl")):
|
||||
raise ValueError("The saved_path must end with .csv or .jsonl.")
|
||||
@@ -173,8 +172,10 @@ def main():
|
||||
min_text_score=args.min_text_score,
|
||||
motion_score_metadata_path=args.motion_score_metadata_path,
|
||||
min_motion_score=args.min_motion_score,
|
||||
semantic_consistency_score_metadata_path=args.semantic_consistency_score_metadata_path,
|
||||
min_semantic_consistency_score=args.min_semantic_consistency_score,
|
||||
video_path_column=args.video_path_column
|
||||
)
|
||||
video_path_list = [os.path.join(args.video_folder, video_path) for video_path in video_path_list]
|
||||
# Sorting to guarantee the same result for each process.
|
||||
video_path_list = natsorted(video_path_list)
|
||||
|
||||
|
||||
@@ -49,6 +49,7 @@ python caption_rewrite.py \
|
||||
--model_name $REWRITE_MODEL_PATH \
|
||||
--prompt prompt/rewrite.txt \
|
||||
--prefix '"rewritten description": ' \
|
||||
--max_retry_count 10 \
|
||||
--saved_path $REWRITTEN_VIDEO_CAPTION_SAVED_PATH \
|
||||
--saved_freq 1
|
||||
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import ast
|
||||
import os
|
||||
from typing import Optional
|
||||
|
||||
import pandas as pd
|
||||
@@ -7,6 +6,7 @@ import pandas as pd
|
||||
from .logger import logger
|
||||
|
||||
|
||||
# Ensure each item in the video_path_list matches the paths in the video_path column of the metadata.
|
||||
def filter(
|
||||
video_path_list: list[str],
|
||||
basic_metadata_path: Optional[str] = None,
|
||||
@@ -28,8 +28,6 @@ def filter(
|
||||
min_semantic_consistency_score: float = 0.80,
|
||||
video_path_column: str = "video_path"
|
||||
):
|
||||
video_path_list = [os.path.basename(video_path) for video_path in video_path_list]
|
||||
|
||||
if basic_metadata_path is not None:
|
||||
if basic_metadata_path.endswith(".csv"):
|
||||
basic_df = pd.read_csv(basic_metadata_path)
|
||||
@@ -39,7 +37,6 @@ def filter(
|
||||
basic_df["resolution"] = basic_df["frame_size"].apply(lambda x: x[0] * x[1])
|
||||
filtered_basic_df = basic_df[basic_df["resolution"] < min_resolution]
|
||||
filtered_video_path_list = filtered_basic_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
@@ -50,7 +47,6 @@ def filter(
|
||||
if min_duration != -1:
|
||||
filtered_basic_df = basic_df[basic_df["duration"] < min_duration]
|
||||
filtered_video_path_list = filtered_basic_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
@@ -61,7 +57,6 @@ def filter(
|
||||
if max_duration != -1:
|
||||
filtered_basic_df = basic_df[basic_df["duration"] > max_duration]
|
||||
filtered_video_path_list = filtered_basic_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
@@ -82,7 +77,6 @@ def filter(
|
||||
aesthetic_score_df["aesthetic_score_mean"] = aesthetic_score_df["aesthetic_score"].apply(lambda x: sum(x) / len(x))
|
||||
filtered_aesthetic_score_df = aesthetic_score_df[aesthetic_score_df["aesthetic_score_mean"] < min_aesthetic_score]
|
||||
filtered_video_path_list = filtered_aesthetic_score_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
@@ -107,7 +101,6 @@ def filter(
|
||||
aesthetic_score_siglip_df["aesthetic_score_siglip_mean"] < min_aesthetic_score_siglip
|
||||
]
|
||||
filtered_video_path_list = filtered_aesthetic_score_siglip_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
@@ -123,7 +116,6 @@ def filter(
|
||||
|
||||
filtered_text_score_df = text_score_df[text_score_df["text_score"] > min_text_score]
|
||||
filtered_video_path_list = filtered_text_score_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
@@ -139,7 +131,6 @@ def filter(
|
||||
|
||||
filtered_motion_score_df = motion_score_df[motion_score_df["motion_score"] < min_motion_score]
|
||||
filtered_video_path_list = filtered_motion_score_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
@@ -149,7 +140,6 @@ def filter(
|
||||
|
||||
filtered_motion_score_df = motion_score_df[motion_score_df["motion_score"] > max_motion_score]
|
||||
filtered_video_path_list = filtered_motion_score_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
@@ -165,7 +155,6 @@ def filter(
|
||||
|
||||
filtered_videoclipxl_score_df = videoclipxl_score_df[videoclipxl_score_df["videoclipxl_score"] < min_videoclipxl_score]
|
||||
filtered_video_path_list = filtered_videoclipxl_score_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
@@ -183,7 +172,6 @@ def filter(
|
||||
semantic_consistency_score_df["similarity_cross_frame"].apply(lambda x: min(x) < min_semantic_consistency_score)
|
||||
]
|
||||
filtered_video_path_list = filtered_semantic_consistency_score_df[video_path_column].tolist()
|
||||
filtered_video_path_list = [os.path.basename(video_path) for video_path in filtered_video_path_list]
|
||||
|
||||
video_path_list = list(set(video_path_list).difference(set(filtered_video_path_list)))
|
||||
logger.info(
|
||||
|
||||
@@ -1,12 +1,13 @@
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
from multiprocessing import Manager, Pool
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
from natsort import index_natsorted
|
||||
|
||||
from .get_meta_file import parallel_rglob
|
||||
from .logger import logger
|
||||
|
||||
|
||||
@@ -20,9 +21,10 @@ def process_file(file_path, shared_list):
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Gather all jsonl files in a folder (meta_folder) to a single jsonl file (meta_file_path).")
|
||||
parser.add_argument("--meta_folder", type=str, required=True)
|
||||
parser.add_argument("--meta_file_path", type=str, required=True)
|
||||
parser.add_argument("--video_path_column", type=str, default="video_path")
|
||||
parser.add_argument("--meta_file_path", type=str, required=True)
|
||||
parser.add_argument("--n_jobs", type=int, default=1)
|
||||
parser.add_argument("--recursive", action="store_true", help="Whether to search sub-folders recursively.")
|
||||
|
||||
args = parser.parse_args()
|
||||
return args
|
||||
@@ -31,7 +33,13 @@ def parse_args():
|
||||
def main():
|
||||
args = parse_args()
|
||||
|
||||
jsonl_files = glob.glob(os.path.join(args.meta_folder, "*.jsonl"))
|
||||
if not os.path.exists(args.meta_folder):
|
||||
raise ValueError(f"The meta_folder {args.meta_folder} does not exist.")
|
||||
meta_folder = Path(args.meta_folder)
|
||||
if args.recursive:
|
||||
jsonl_files = [str(file) for file in parallel_rglob(meta_folder, f"*.jsonl", max_workers=args.n_jobs)]
|
||||
else:
|
||||
jsonl_files = [str(file) for file in meta_folder.glob(f"*.jsonl")]
|
||||
|
||||
with Manager() as manager:
|
||||
shared_list = manager.list()
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
import argparse
|
||||
import os
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
|
||||
import pandas as pd
|
||||
@@ -7,9 +9,23 @@ from tqdm import tqdm
|
||||
|
||||
from .logger import logger
|
||||
|
||||
ALL_VIDEO_EXT = set(["mp4", "webm", "mkv", "avi", "flv", "mov"])
|
||||
ALL_VIDEO_EXT = set(["mp4", "webm", "mkv", "avi", "flv", "mov", "rmvb"])
|
||||
ALL_IMGAE_EXT = set(["png", "webp", "jpg", "jpeg", "bmp", "gif"])
|
||||
|
||||
def parallel_rglob(root_path, pattern, max_workers=8):
|
||||
root = Path(root_path)
|
||||
futures = []
|
||||
results = []
|
||||
with ThreadPoolExecutor(max_workers=max_workers) as executor:
|
||||
for sub_path in root.iterdir():
|
||||
if sub_path.is_dir():
|
||||
futures.append(executor.submit(lambda p=sub_path: list(p.rglob(pattern))))
|
||||
for future in as_completed(futures):
|
||||
results.extend(future.result())
|
||||
results.extend(root.glob(pattern))
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(description="Compute scores of uniform sampled frames from videos.")
|
||||
@@ -41,6 +57,10 @@ def main():
|
||||
raise ValueError("Either video_folder or image_folder should be specified in the arguments.")
|
||||
if args.video_folder is not None and args.image_folder is not None:
|
||||
raise ValueError("Both video_folder and image_folder can not be specified in the arguments at the same time.")
|
||||
if args.image_folder is None and not os.path.exists(args.video_folder):
|
||||
raise ValueError(f"The video_folder {args.video_folder} does not exist.")
|
||||
if args.video_folder is None and not os.path.exists(args.image_folder):
|
||||
raise ValueError(f"The image_folder {args.image_folder} does not exist.")
|
||||
|
||||
# Use the path name instead of the file name as video_path/image_path (unique ID).
|
||||
if args.video_folder is not None:
|
||||
@@ -48,7 +68,7 @@ def main():
|
||||
video_folder = Path(args.video_folder)
|
||||
for ext in tqdm(list(ALL_VIDEO_EXT)):
|
||||
if args.recursive:
|
||||
video_path_list += [str(file.relative_to(video_folder)) for file in video_folder.rglob(f"*.{ext}")]
|
||||
video_path_list += [str(file.relative_to(video_folder)) for file in parallel_rglob(video_folder, f"*.{ext}")]
|
||||
else:
|
||||
video_path_list += [str(file.relative_to(video_folder)) for file in video_folder.glob(f"*.{ext}")]
|
||||
video_path_list = natsorted(video_path_list)
|
||||
@@ -59,7 +79,7 @@ def main():
|
||||
image_folder = Path(args.image_folder)
|
||||
for ext in tqdm(list(ALL_IMGAE_EXT)):
|
||||
if args.recursive:
|
||||
image_path_list += [str(file.relative_to(image_folder)) for file in image_folder.rglob(f"*.{ext}")]
|
||||
image_path_list += [str(file.relative_to(image_folder)) for file in parallel_rglob(video_folder, f"*.{ext}")]
|
||||
else:
|
||||
image_path_list += [str(file.relative_to(image_folder)) for file in image_folder.glob(f"*.{ext}")]
|
||||
image_path_list = natsorted(image_path_list)
|
||||
|
||||
@@ -41,7 +41,8 @@ def clip_video(video_path, timecode_list, output_folder, video_duration):
|
||||
according to the timecode obtained from easyanimate/video_caption/cutscene_detect.py.
|
||||
"""
|
||||
try:
|
||||
video_name = Path(video_path).stem
|
||||
os.makedirs(output_folder, exist_ok=True)
|
||||
video_stem = Path(video_path).stem
|
||||
|
||||
if len(timecode_list) == 0: # The video of a single scene.
|
||||
splitted_timecode_list = []
|
||||
@@ -57,7 +58,7 @@ def clip_video(video_path, timecode_list, output_folder, video_duration):
|
||||
splitted_index += 1
|
||||
continue
|
||||
splitted_timecode_list.append([cur_start.strftime("%H:%M:%S.%f")[:-3], cur_end.strftime("%H:%M:%S.%f")[:-3]])
|
||||
output_path = os.path.join(output_folder, video_name + f"_{splitted_index}.mp4")
|
||||
output_path = os.path.join(output_folder, video_stem + f"_{splitted_index}.mp4")
|
||||
if os.path.exists(output_path):
|
||||
logger.info(f"The clipped video {output_path} exists.")
|
||||
cur_start = cur_end
|
||||
@@ -77,7 +78,7 @@ def clip_video(video_path, timecode_list, output_folder, video_duration):
|
||||
start_time = datetime.strptime(timecode[0], "%H:%M:%S.%f")
|
||||
end_time = datetime.strptime(timecode[1], "%H:%M:%S.%f")
|
||||
video_duration = (end_time - start_time).total_seconds()
|
||||
output_path = os.path.join(output_folder, video_name + f"_{i}.mp4")
|
||||
output_path = os.path.join(output_folder, video_stem + f"_{i}.mp4")
|
||||
if os.path.exists(output_path):
|
||||
logger.info(f"The clipped video {output_path} exists.")
|
||||
continue
|
||||
@@ -93,7 +94,7 @@ def clip_video(video_path, timecode_list, output_folder, video_duration):
|
||||
if cur_video_duration < MIN_SECONDS:
|
||||
break
|
||||
splitted_timecode_list.append([cur_start.strftime("%H:%M:%S.%f")[:-3], cur_end.strftime("%H:%M:%S.%f")[:-3]])
|
||||
splitted_output_path = os.path.join(output_folder, video_name + f"_{i}_{splitted_index}.mp4")
|
||||
splitted_output_path = os.path.join(output_folder, video_stem + f"_{i}_{splitted_index}.mp4")
|
||||
if os.path.exists(splitted_output_path):
|
||||
logger.info(f"The clipped video {splitted_output_path} exists.")
|
||||
cur_start = cur_end
|
||||
@@ -145,19 +146,23 @@ if __name__ == "__main__":
|
||||
video_metadata_df = video_metadata_df[video_metadata_df["resolution"] >= args.resolution_threshold]
|
||||
logger.info(f"Filter {num_videos - len(video_metadata_df)} videos with resolution smaller than {args.resolution_threshold}.")
|
||||
video_path_list = video_metadata_df[args.video_path_column].to_list()
|
||||
video_id_list = [Path(video_path).stem for video_path in video_path_list]
|
||||
if len(video_id_list) != len(list(set(video_id_list))):
|
||||
logger.warning("Duplicate file names exist in the input video path list.")
|
||||
video_path_list = [os.path.join(args.video_folder, video_path) for video_path in video_path_list]
|
||||
video_timecode_list = video_metadata_df["timecode_list"].to_list()
|
||||
video_duration_list = video_metadata_df["duration"].to_list()
|
||||
|
||||
assert len(video_path_list) == len(video_timecode_list)
|
||||
os.makedirs(args.output_folder, exist_ok=True)
|
||||
if args.video_folder == "":
|
||||
output_folder_list = [args.output_folder] * len(video_path_list)
|
||||
video_name_list = [Path(video_path).name for video_path in video_path_list]
|
||||
# We only check the unique video name with the absolute video path.
|
||||
if len(video_name_list) != len(set(video_name_list)):
|
||||
logger.error(f"The video path in {args.video_metadata_path} should has an unique video name.")
|
||||
else:
|
||||
output_folder_list = [os.path.join(args.output_folder, os.path.dirname(video_path)) for video_path in video_path_list]
|
||||
video_path_list = [os.path.join(args.video_folder, video_path) for video_path in video_path_list]
|
||||
|
||||
args_list = [
|
||||
(video_path, timecode_list, args.output_folder, video_duration)
|
||||
for video_path, timecode_list, video_duration in zip(
|
||||
video_path_list, video_timecode_list, video_duration_list
|
||||
(video_path, timecode_list, output_folder, video_duration)
|
||||
for video_path, timecode_list, output_folder, video_duration in zip(
|
||||
video_path_list, video_timecode_list, output_folder_list, video_duration_list
|
||||
)
|
||||
]
|
||||
with Pool(args.n_jobs) as pool:
|
||||
|
||||
@@ -1177,13 +1177,12 @@ def main():
|
||||
)
|
||||
|
||||
for epoch in range(first_epoch, args.num_train_epochs):
|
||||
train_dataloader_iterations = 100
|
||||
train_loss = 0.0
|
||||
train_reward = 0.0
|
||||
|
||||
# In the following training loop, randomly select training prompts and use the
|
||||
# `EasyAnimatePipeline_Multi_Text_Encoder_Inpaint` to sample videos, calculate rewards, and update the network.
|
||||
for _ in range(train_dataloader_iterations):
|
||||
for _ in range(num_update_steps_per_epoch):
|
||||
# train_prompt = random.sample(prompt_list, args.train_batch_size)
|
||||
train_prompt = random.choices(prompt_list, k=args.train_batch_size)
|
||||
logger.info(f"train_prompt: {train_prompt}")
|
||||
|
||||
Reference in New Issue
Block a user