| 284 | return self.forward_test(x) |
| 285 | |
| 286 | def forward_test(self, x): |
| 287 | if not self.ds: |
| 288 | visual_feats = x + self.vis_pos_embed |
| 289 | else: |
| 290 | visual_feats = x |
| 291 | bs = visual_feats.shape[0] |
| 292 | |
| 293 | pos_node_embed = self.pos_node_embed( |
| 294 | torch.arange(self.max_len).cuda( |
| 295 | x.get_device())).unsqueeze(0) + self.char_pos_embed |
| 296 | pos_node_embed = torch.tile(pos_node_embed, [bs, 1, 1]) |
| 297 | |
| 298 | char_vis_node_query = visual_feats |
| 299 | pos_vis_node_query = torch.concat([pos_node_embed, visual_feats], 1) |
| 300 | |
| 301 | for char_decoder_layer, pos_decoder_layer in zip( |
| 302 | self.char_node_decoder, self.pos_node_decoder): |
| 303 | char_vis_node_query = char_decoder_layer(char_vis_node_query, |
| 304 | char_vis_node_query) |
| 305 | pos_vis_node_query = pos_decoder_layer( |
| 306 | pos_vis_node_query, pos_vis_node_query[:, self.max_len:, :]) |
| 307 | |
| 308 | pos_node_query = pos_vis_node_query[:, :self.max_len, :] |
| 309 | |
| 310 | char_vis_feats = char_vis_node_query |
| 311 | # pos_vis_feats = pos_vis_node_query[:, self.max_len :, :] |
| 312 | |
| 313 | # pos_node_feats = self.edge_decoder( |
| 314 | # pos_node_query, char_vis_feats, pos_vis_feats |
| 315 | # ) # B, 26, dim |
| 316 | |
| 317 | pos_node_feats = pos_node_query |
| 318 | for layer_i in range(self.rec_layer_num): |
| 319 | rec_layer = self.edge_decoder[layer_i] |
| 320 | if (self.rec_layer_num + layer_i) % 2 == 0: |
| 321 | pos_node_feats = rec_layer(pos_node_feats, pos_node_feats, |
| 322 | self.self_mask) |
| 323 | else: |
| 324 | pos_node_feats = rec_layer(pos_node_feats, char_vis_feats) |
| 325 | edge_feats = self.edge_fc(pos_node_feats) # B, 26, 37 |
| 326 | |
| 327 | edge_logits = F.softmax( |
| 328 | edge_feats, |
| 329 | -1) # * F.sigmoid(pos_node_feats1.unsqueeze(-1)) # B, 26, 37 |
| 330 | |
| 331 | return edge_logits |
| 332 | |
| 333 | def forward_train(self, x, targets=None): |
| 334 | if not self.ds: |