From 2a97f7632b659b254d1d77de150d0f78ebba64ca Mon Sep 17 00:00:00 2001 From: hkunzhe Date: Tue, 21 Jan 2025 20:10:28 +0800 Subject: [PATCH] fix extract frames in compute_semantic_consistency --- .../compute_semantic_consistency.py | 16 ++++++++++--- .../video_caption/utils/video_utils.py | 24 ++++++++++++------- 2 files changed, 29 insertions(+), 11 deletions(-) diff --git a/easyanimate/video_caption/compute_semantic_consistency.py b/easyanimate/video_caption/compute_semantic_consistency.py index c12bf23..4599cfe 100644 --- a/easyanimate/video_caption/compute_semantic_consistency.py +++ b/easyanimate/video_caption/compute_semantic_consistency.py @@ -170,8 +170,18 @@ def main(): for idx, batch in enumerate(tqdm(video_loader)): if len(batch) > 0: - batch_video_path = batch["path"] - batch_frame = batch["sampled_frame"] + batch_video_path = [] + batch_frame = [] + batch_sampled_frame_idx = [] + # At least two frames are required to calculate cross-frame semantic consistency. + for path, frame, frame_idx in zip(batch["path"], batch["sampled_frame"], batch["sampled_frame_idx"]): + if len(frame) > 1: + batch_video_path.append(path) + batch_frame.append(frame) + batch_sampled_frame_idx.append(frame_idx) + else: + logger.warning(f"Skip {path} because it only has {len(frame)} frames.") + frame_num_list = [len(video_frames) for video_frames in batch_frame] # [B, T, H, W, C] => [(B * T), H, W, C] reshaped_batch_frame = [frame for video_frames in batch_frame for frame in video_frames] @@ -197,7 +207,7 @@ def main(): result_dict[args.video_path_column].extend(saved_video_path_list) result_dict["similarity_cross_frame"].extend(batch_simi_cross_frame) result_dict["similarity_mean"].extend(batch_similarity_mean) - result_dict["sample_frame_idx"].extend(batch["sampled_frame_idx"]) + result_dict["sample_frame_idx"].extend(batch_sampled_frame_idx) # Save the metadata in the main process every saved_freq. if (idx % args.saved_freq) == 0 or idx == len(video_loader) - 1: diff --git a/easyanimate/video_caption/utils/video_utils.py b/easyanimate/video_caption/utils/video_utils.py index 2c7e12e..bc9cc6c 100644 --- a/easyanimate/video_caption/utils/video_utils.py +++ b/easyanimate/video_caption/utils/video_utils.py @@ -44,13 +44,16 @@ def get_keyframe_index(video_path): result = subprocess.run(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, universal_newlines=True) keyframe_index_list = [] - for index, line in enumerate(result.stdout.split("\n")): + frame_index = 0 + for line in result.stdout.split("\n"): line = line.strip(",") pict_type = line.strip() if pict_type == "I": - keyframe_index_list.append(index) + keyframe_index_list.append(frame_index) + if pict_type == "I" or pict_type == "B" or pict_type == "P": + frame_index += 1 - return keyframe_index_list + return keyframe_index_list, frame_index def extract_frames( video_path: str, @@ -81,17 +84,22 @@ def extract_frames( elif sample_method == "last": sampled_frame_idx_list = [len(vr) - 1] elif sample_method == "keyframe": - sampled_frame_idx_list = get_keyframe_index(video_path) - elif sample_method == "keyframe+first": - sampled_frame_idx_list = get_keyframe_index(video_path) + sampled_frame_idx_list, final_frame_index = get_keyframe_index(video_path) + elif sample_method == "keyframe+first": # keyframe + the first second + sampled_frame_idx_list, final_frame_index = get_keyframe_index(video_path) if len(sampled_frame_idx_list) == 1 or sampled_frame_idx_list[1] > 1 * vr.get_avg_fps(): + if int(1 * vr.get_avg_fps()) > len(vr): + raise ValueError(f"The duration of {video_path} is less than 1s.") sampled_frame_idx_list.insert(1, int(1 * vr.get_avg_fps())) - elif sample_method == "keyframe+last": - sampled_frame_idx_list = get_keyframe_index(video_path) + elif sample_method == "keyframe+last": # keyframe + the last frame + sampled_frame_idx_list, final_frame_index = get_keyframe_index(video_path) if sampled_frame_idx_list[-1] != (len(vr) - 1): sampled_frame_idx_list.append(len(vr) - 1) else: raise ValueError(f"The sample_method must be within {ALL_FRAME_SAMPLE_METHODS}.") + if "keyframe" in sample_method: + if final_frame_index != len(vr): + raise ValueError(f"The keyframe index list is not accurate. Please check the video {video_path}.") sampled_frame_list = vr.get_batch(sampled_frame_idx_list).asnumpy() sampled_frame_list = [Image.fromarray(frame) for frame in sampled_frame_list]