Fix train reward lora and video caption (#165)

This commit is contained in:
hkz
2024-12-10 14:23:33 +08:00
committed by GitHub
parent ccbdce4492
commit 485e771acf
15 changed files with 172 additions and 83 deletions
@@ -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)
+1
View File
@@ -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
```
+61 -20
View File
@@ -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}.")
+8 -8
View File
@@ -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 -13
View File
@@ -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)
+18 -13
View File
@@ -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:
+1 -2
View File
@@ -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}")