fix extract frames in compute_semantic_consistency
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
Reference in New Issue
Block a user