Spaces:
Sleeping
Sleeping
| 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) | |