import torch import math import random import copy def get_mask_from_mask_type(video_length, mask_type, mask_args): if mask_type == "t2v": # 文生视频 mask = torch.zeros(video_length, dtype=torch.int64) elif mask_type == "frame_interpolation": # 传统插帧 frame_interpolation_rate = mask_args["frame_interpolation_rate"] a = video_length // frame_interpolation_rate + 1 mask = [torch.ones(a, dtype=torch.int64)] for i in range(frame_interpolation_rate - 1): mask.append(torch.zeros(a, dtype=torch.int64)) mask = torch.stack(mask, dim=-1).reshape(-1)[:video_length] elif mask_type == "image_interpolation": # 两图片插帧 mask = torch.zeros(video_length, dtype=torch.int64) mask[0] = 1 mask[-1] = 1 elif mask_type == "video_connect": # 两视频插帧 mask = torch.cat((torch.ones(video_length//4, dtype=torch.int64), torch.zeros(video_length-((video_length//4)*2), dtype=torch.int64), torch.ones(video_length//4, dtype=torch.int64))) elif mask_type == "random": # 随机mask mask = torch.randint(0, 2, (video_length,), dtype=torch.int64).float() elif mask_type == "expansion": # 视频扩展 expansion_save_rate = mask_args["expansion_save_rate"] expansion_reverse_rate = mask_args["expansion_reverse_rate"] random_num = random.random() # [0, 1) if random_num < expansion_reverse_rate: expansion_reverse = True else: expansion_reverse = False video_length_save = math.floor(video_length * expansion_save_rate) video_length_drop = video_length - video_length_save assert video_length_drop != 0, f"video_length_drop: {video_length_drop} must be non-zero." mask_save = torch.ones((video_length_save), dtype=torch.int64) mask_drop = torch.zeros((video_length_drop), dtype=torch.int64) if expansion_reverse: mask = torch.concat([mask_drop, mask_save], dim=-1) else: mask = torch.concat([mask_save, mask_drop], dim=-1) elif mask_type == "i2v": # 图生视频 mask = torch.zeros(video_length, dtype=torch.int64) mask[0] = 1 else: raise ValueError(f"{mask_type} is not supported !") mask = mask.to(torch.int64) return mask def multi_task_check(multi_task_args): count_prob = 0 for k, v in multi_task_args.items(): assert k in ["t2v", "frame_interpolation", "image_interpolation", "expansion", "i2v", "video_connect", "random"], \ f"{k} must be within [t2v, frame_interpolation, image_interpolation, expansion, i2v, video_connect, random]" count_prob += v["prob"] if k == "frame_interpolation": assert "mask_args" in v, f"mask_args must be provided for frame_interpolation" assert "frame_interpolation_rate" in v["mask_args"], f"frame_interpolation_rate must be provided for frame_interpolation" elif k == "expansion": assert "mask_args" in v, f"mask_args must be provided for expansion" assert "expansion_save_rate" in v["mask_args"], f"expansion_save_rate must be provided for expansion" assert "expansion_reverse_rate" in v["mask_args"], f"expansion_reverse_rate must be provided for expansion" assert abs(1 - count_prob) < 0.000001, f"sum(prob): {count_prob} must be 1" def get_multi_task_mask(video_length, multi_task_args): cumsum_threshold = 0 thresholds = [0] mask_types = [] for k, v in multi_task_args.items(): cumsum_threshold += v["prob"] thresholds.append(cumsum_threshold) mask_types.append(k) random_num = random.random() for i in range(0, len(thresholds) - 1): if thresholds[i] <= random_num < thresholds[i + 1]: mask_type = mask_types[i] if mask_type == "frame_interpolation" or mask_type == "expansion": mask_args = multi_task_args[mask_type]["mask_args"] else: mask_args = None break multi_task_mask = get_mask_from_mask_type(video_length, mask_type=mask_type, mask_args=mask_args) return multi_task_mask, mask_type def merge_tensor_by_mask(tensor_1, tensor_2, mask, dim): assert tensor_1.shape == tensor_2.shape # Mask is a 0/1 verctor. Choose tensor_2, when the value is 1; otherwise, tensor_1 masked_indices = torch.nonzero(mask).squeeze(1) tmp = copy.deepcopy(tensor_1) if dim==0: tmp[masked_indices] = tensor_2[masked_indices] if dim==1: tmp[:, masked_indices] = tensor_2[:, masked_indices] if dim==2: tmp[:, :, masked_indices] = tensor_2[:, :, masked_indices] # print(f'mask-{mask}-masked_indices-{masked_indices}-{tensor_2[:, :, masked_indices].shape}-{tensor_2[:, :, masked_indices]}-{tmp[:, :, masked_indices]}') return tmp if __name__ == '__main__': video_length = 129 # multi_task_args = {'i2v': {'prob': 1}} # multi_task_args = {'expansion': {'prob': 1.0, "mask_args": {"expansion_save_rate": 0.5, "expansion_reverse_rate": 0.0}}} multi_task_args = {'random': {'prob': 1.0}} multitask_mask, mask_type = get_multi_task_mask(video_length, multi_task_args) print('multitask_mask = ', multitask_mask) print('mask_type = ', mask_type)