(self, args)
| 10 | |
| 11 | class MULT(nn.Module): |
| 12 | def __init__(self, args): |
| 13 | super(MULT, self).__init__() |
| 14 | |
| 15 | # params: analyze args |
| 16 | audio_dim = 1024 # args.audio_dim |
| 17 | text_dim = 4096 # args.text_dim |
| 18 | video_dim = 1024 # args.video_dim |
| 19 | output_dim1 = args.output_dim1 |
| 20 | output_dim2 = args.output_dim2 |
| 21 | |
| 22 | # params: analyze args |
| 23 | self.attn_mask = True |
| 24 | self.layers = 4# args.layers # 4 |
| 25 | self.dropout = args.dropout |
| 26 | self.num_heads = 2 # args.num_heads # 8 |
| 27 | self.hidden_dim = 128 # args.hidden_dim # 128 |
| 28 | self.conv1d_kernel_size = 5 #args.conv1d_kernel_size # 5 |
| 29 | # self.grad_clip = args.grad_clip |
| 30 | |
| 31 | # params: intermedia |
| 32 | combined_dim = 2 * (self.hidden_dim + self.hidden_dim + self.hidden_dim) |
| 33 | output_dim = self.hidden_dim // 2 |
| 34 | |
| 35 | # 1. Temporal convolutional layers |
| 36 | self.proj_l = nn.Conv1d(text_dim, self.hidden_dim, kernel_size=self.conv1d_kernel_size, padding=0, bias=False) |
| 37 | self.proj_a = nn.Conv1d(audio_dim, self.hidden_dim, kernel_size=self.conv1d_kernel_size, padding=0, bias=False) |
| 38 | self.proj_v = nn.Conv1d(video_dim, self.hidden_dim, kernel_size=self.conv1d_kernel_size, padding=0, bias=False) |
| 39 | |
| 40 | # 2. Crossmodal Attentions |
| 41 | self.trans_l_with_a = self.get_network(self_type='la') |
| 42 | self.trans_l_with_v = self.get_network(self_type='lv') |
| 43 | |
| 44 | self.trans_a_with_l = self.get_network(self_type='al') |
| 45 | self.trans_a_with_v = self.get_network(self_type='av') |
| 46 | |
| 47 | self.trans_v_with_l = self.get_network(self_type='vl') |
| 48 | self.trans_v_with_a = self.get_network(self_type='va') |
| 49 | |
| 50 | # 3. Self Attentions (Could be replaced by LSTMs, GRUs, etc.) |
| 51 | self.trans_l_mem = self.get_network(self_type='l_mem', layers=3) |
| 52 | self.trans_a_mem = self.get_network(self_type='a_mem', layers=3) |
| 53 | self.trans_v_mem = self.get_network(self_type='v_mem', layers=3) |
| 54 | |
| 55 | # Projection layers |
| 56 | self.proj1 = nn.Linear(combined_dim, combined_dim) |
| 57 | self.proj2 = nn.Linear(combined_dim, combined_dim) |
| 58 | self.out_layer = nn.Linear(combined_dim, output_dim) |
| 59 | |
| 60 | # cls layers |
| 61 | self.fc_out_1 = nn.Linear(output_dim, output_dim1) |
| 62 | self.fc_out_2 = nn.Linear(output_dim, output_dim2) |
| 63 | |
| 64 | |
| 65 | def get_network(self, self_type='l', layers=-1): |
nothing calls this directly
no test coverage detected