audio_feat: tensor of shape (batch, seqlen1, audio_in) video_feat: tensor of shape (batch, seqlen2, video_in) text_feat: tensor of shape (batch, seqlen3, text_in)
(self, batch)
| 90 | |
| 91 | |
| 92 | def forward(self, batch): |
| 93 | ''' |
| 94 | audio_feat: tensor of shape (batch, seqlen1, audio_in) |
| 95 | video_feat: tensor of shape (batch, seqlen2, video_in) |
| 96 | text_feat: tensor of shape (batch, seqlen3, text_in) |
| 97 | ''' |
| 98 | # x_l = batch['texts'].transpose(1, 2) |
| 99 | # x_a = batch['audios'].transpose(1, 2) |
| 100 | # x_v = batch['videos'].transpose(1, 2) |
| 101 | x_l = batch[1].transpose(1, 2) |
| 102 | x_a = batch[0].transpose(1, 2) |
| 103 | x_v = batch[2].transpose(1, 2) |
| 104 | |
| 105 | # Project the textual/visual/audio features |
| 106 | proj_x_l = self.proj_l(x_l).permute(2, 0, 1) |
| 107 | proj_x_a = self.proj_a(x_a).permute(2, 0, 1) |
| 108 | proj_x_v = self.proj_v(x_v).permute(2, 0, 1) |
| 109 | |
| 110 | # (V,A) --> L |
| 111 | h_l_with_as = self.trans_l_with_a(proj_x_l, proj_x_a, proj_x_a) |
| 112 | h_l_with_vs = self.trans_l_with_v(proj_x_l, proj_x_v, proj_x_v) |
| 113 | h_ls = torch.cat([h_l_with_as, h_l_with_vs], dim=2) |
| 114 | h_ls = self.trans_l_mem(h_ls) |
| 115 | if type(h_ls) == tuple: |
| 116 | h_ls = h_ls[0] |
| 117 | last_h_l = last_hs = h_ls[-1] |
| 118 | |
| 119 | # (L,V) --> A |
| 120 | h_a_with_ls = self.trans_a_with_l(proj_x_a, proj_x_l, proj_x_l) |
| 121 | h_a_with_vs = self.trans_a_with_v(proj_x_a, proj_x_v, proj_x_v) |
| 122 | h_as = torch.cat([h_a_with_ls, h_a_with_vs], dim=2) |
| 123 | h_as = self.trans_a_mem(h_as) |
| 124 | if type(h_as) == tuple: |
| 125 | h_as = h_as[0] |
| 126 | last_h_a = last_hs = h_as[-1] |
| 127 | |
| 128 | # (L,A) --> V |
| 129 | h_v_with_ls = self.trans_v_with_l(proj_x_v, proj_x_l, proj_x_l) |
| 130 | h_v_with_as = self.trans_v_with_a(proj_x_v, proj_x_a, proj_x_a) |
| 131 | h_vs = torch.cat([h_v_with_ls, h_v_with_as], dim=2) |
| 132 | h_vs = self.trans_v_mem(h_vs) |
| 133 | if type(h_vs) == tuple: |
| 134 | h_vs = h_vs[0] |
| 135 | last_h_v = last_hs = h_vs[-1] |
| 136 | last_hs = torch.cat([last_h_l, last_h_a, last_h_v], dim=1) |
| 137 | |
| 138 | # A residual block |
| 139 | last_hs_proj = self.proj2(F.dropout(F.relu(self.proj1(last_hs), inplace=True), p=self.dropout, training=self.training)) |
| 140 | last_hs_proj += last_hs |
| 141 | features = self.out_layer(last_hs_proj) |
| 142 | |
| 143 | # store results |
| 144 | # emos_out = self.fc_out_1(features) |
| 145 | vals_out = self.fc_out_2(features) |
| 146 | interloss = torch.tensor(0).cuda() |
| 147 | |
| 148 | return features, vals_out, interloss |
nothing calls this directly
no outgoing calls
no test coverage detected