File size: 3,099 Bytes
5de1792
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import torch
import torch.nn as nn
import numpy as np
import math
import random
#temoral-frame mask
class TemporalMask:
    def __init__(self, mask_ratio, person_num, joint_num, channel_num):
        self.mask_ratio = mask_ratio
        self.person_num = person_num
        self.joint_num = joint_num
        self.channel_num = channel_num
        
    def __call__(self, data, frame):
        '''
        given a data: N,C,T,V,M
                frame: the number of the valid frames in data 
        return Mask: N,C,T,V,M 
        randomly mask the frame*ratio frames
        '''
        #data N,C,T,V,M
        N,C,T,V,M = data.shape
        data = data.permute((0,2,4,3,1))##N,T,M,V,C
        #data = data.reshape(N,T,C*V*M)
        size, max_frame, feature_dim = N,T,V*C*M
        #data = data.view(size, max_frame, self.person_num, self.joint_num, self.channel_num)#N,T,M,V,C

        mask_idx = torch.tensor((frame * (1 - self.mask_ratio))).reshape((1, 1, 1, 1, 1)).repeat(size, max_frame, self.person_num, self.joint_num, self.channel_num)
        frame_idx = torch.arange(max_frame).reshape((1, max_frame, 1, 1, 1)).repeat(size, 1, self.person_num, self.joint_num, self.channel_num)
        mask = (frame_idx < mask_idx).float()
        randper = np.random.permutation(range(frame))
        rand = torch.arange(max_frame)
        rand[0:frame] = torch.from_numpy(randper)
        mask = mask[:,rand,...]
        #N,T,M,V,C -> N C T V M
        #mask = mask.permute((0,4,1,3,2))
        #trans_data = data * mask
        #trans_data = trans_data.view(size, max_frame, feature_dim)
        return mask #N,T,M,V,C
    
#random joint mask different for frames
class Jointmask:
    def __init__(self, mask_ratio, person_num, joint_num, channel_num):
        self.mask_ratio = mask_ratio
        self.person_num = person_num
        self.joint_num = joint_num
        self.channel_num = channel_num
    def __call__(self, data, frame):
        N,C,T,V,M = data.shape
        data = data.permute((0,2,4,3,1))##N,T,M,V,C
        mask = torch.ones((T,V))
        mask_joint_num = int(V*self.mask_ratio)
        for i in range(frame):
            rand = random.sample(range(0,V),mask_joint_num)
            mask[i][rand] = 0.0
        mask = mask.reshape(1,T,1,V,1).repeat(N,1,M,1,C)
        return mask

#random joint mask same for frames
class Jointmask2:
    def __init__(self, mask_ratio, person_num, joint_num, channel_num):
        self.mask_ratio = mask_ratio
        self.person_num = person_num
        self.joint_num = joint_num
        self.channel_num = channel_num
    def __call__(self, data, frame):
        N,C,T,V,M = data.shape
        data = data.permute((0,2,4,3,1))##N,T,M,V,C
        mask = torch.ones((N,T,V))
        mask_joint_num = int(V*self.mask_ratio)
        rand = random.sample(range(0,V),mask_joint_num)
        mask[:,:frame,rand] = 0.0
        mask = mask.reshape(N,T,1,V,1).repeat(1,1,M,1,C)
        return mask


if __name__ =='__main__':
    J = Jointmask(0.5,1,10,1)
    T = TemporalMask(0.5,1,10,1)
    mask = J(torch.rand((1,1,3,10,1)),3)
    print(mask)