File size: 5,435 Bytes
d3ea518
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
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)