homa / hyavatar /utils /multi_task_utils.py
multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
d3ea518 verified
Raw
History Blame Contribute Delete
5.44 kB
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)