MCPcopy Create free account
hub / github.com/WarmCongee/SDUMC / forward

Method forward

toolkit/models/mult.py:92–148  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected