diff --git a/easyanimate/reward/MPS/trainer/models/cross_modeling.py b/easyanimate/reward/MPS/trainer/models/cross_modeling.py index 6822329..31dcee5 100644 --- a/easyanimate/reward/MPS/trainer/models/cross_modeling.py +++ b/easyanimate/reward/MPS/trainer/models/cross_modeling.py @@ -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) diff --git a/easyanimate/video_caption/README.md b/easyanimate/video_caption/README.md index 54cbb4d..5bc5be5 100644 --- a/easyanimate/video_caption/README.md +++ b/easyanimate/video_caption/README.md @@ -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 ``` diff --git a/easyanimate/video_caption/README_zh-CN.md b/easyanimate/video_caption/README_zh-CN.md index dcf0e32..7b12181 100644 --- a/easyanimate/video_caption/README_zh-CN.md +++ b/easyanimate/video_caption/README_zh-CN.md @@ -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 ``` diff --git a/easyanimate/video_caption/caption_rewrite.py b/easyanimate/video_caption/caption_rewrite.py index 3f50bfd..c70fc41 100644 --- a/easyanimate/video_caption/caption_rewrite.py +++ b/easyanimate/video_caption/caption_rewrite.py @@ -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): diff --git a/easyanimate/video_caption/compute_motion_score.py b/easyanimate/video_caption/compute_motion_score.py index d327dd8..b54f7b3 100644 --- a/easyanimate/video_caption/compute_motion_score.py +++ b/easyanimate/video_caption/compute_motion_score.py @@ -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 diff --git a/easyanimate/video_caption/compute_semantic_consistency.py b/easyanimate/video_caption/compute_semantic_consistency.py index 6e1c626..c12bf23 100644 --- a/easyanimate/video_caption/compute_semantic_consistency.py +++ b/easyanimate/video_caption/compute_semantic_consistency.py @@ -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: diff --git a/easyanimate/video_caption/compute_text_score.py b/easyanimate/video_caption/compute_text_score.py index 8b73afd..3acebe1 100644 --- a/easyanimate/video_caption/compute_text_score.py +++ b/easyanimate/video_caption/compute_text_score.py @@ -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}.") diff --git a/easyanimate/video_caption/cutscene_detect.py b/easyanimate/video_caption/cutscene_detect.py index 7ba0014..4e45110 100644 --- a/easyanimate/video_caption/cutscene_detect.py +++ b/easyanimate/video_caption/cutscene_detect.py @@ -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. diff --git a/easyanimate/video_caption/internvl2_video_recaptioning.py b/easyanimate/video_caption/internvl2_video_recaptioning.py index 023ad2b..f779731 100644 --- a/easyanimate/video_caption/internvl2_video_recaptioning.py +++ b/easyanimate/video_caption/internvl2_video_recaptioning.py @@ -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) diff --git a/easyanimate/video_caption/scripts/stage_3_video_recaptioning.sh b/easyanimate/video_caption/scripts/stage_3_video_recaptioning.sh index f0d3451..9f0e896 100644 --- a/easyanimate/video_caption/scripts/stage_3_video_recaptioning.sh +++ b/easyanimate/video_caption/scripts/stage_3_video_recaptioning.sh @@ -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 diff --git a/easyanimate/video_caption/utils/filter.py b/easyanimate/video_caption/utils/filter.py index d9571a4..f7a0c59 100644 --- a/easyanimate/video_caption/utils/filter.py +++ b/easyanimate/video_caption/utils/filter.py @@ -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( diff --git a/easyanimate/video_caption/utils/gather_jsonl.py b/easyanimate/video_caption/utils/gather_jsonl.py index bc3836c..e753053 100644 --- a/easyanimate/video_caption/utils/gather_jsonl.py +++ b/easyanimate/video_caption/utils/gather_jsonl.py @@ -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() diff --git a/easyanimate/video_caption/utils/get_meta_file.py b/easyanimate/video_caption/utils/get_meta_file.py index 85c973f..25368b4 100644 --- a/easyanimate/video_caption/utils/get_meta_file.py +++ b/easyanimate/video_caption/utils/get_meta_file.py @@ -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) diff --git a/easyanimate/video_caption/video_splitting.py b/easyanimate/video_caption/video_splitting.py index 8b2ccab..1241522 100644 --- a/easyanimate/video_caption/video_splitting.py +++ b/easyanimate/video_caption/video_splitting.py @@ -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: diff --git a/scripts/train_reward_lora.py b/scripts/train_reward_lora.py index 5ab7d83..ac4ec3c 100644 --- a/scripts/train_reward_lora.py +++ b/scripts/train_reward_lora.py @@ -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}")