Args: tokens: (N, T, C) where T is length, N is batch size and C is classes number lengths: (N,)
(self, img_feat, data=None)
| 76 | return logits |
| 77 | |
| 78 | def forward(self, img_feat, data=None): |
| 79 | """ |
| 80 | Args: |
| 81 | tokens: (N, T, C) where T is length, N is batch size and C is classes number |
| 82 | lengths: (N,) |
| 83 | """ |
| 84 | img_feat = img_feat + self.v_embeding |
| 85 | B, L, C = img_feat.shape |
| 86 | |
| 87 | # -------------------------------------------------------------------------- |
| 88 | # decoder procedure |
| 89 | T = self.max_length |
| 90 | zeros = img_feat.new_zeros((B, T, C)) |
| 91 | zeros_len = img_feat.new_zeros(B) |
| 92 | query = self.pos_encoder(zeros) |
| 93 | |
| 94 | # 1. vision decode |
| 95 | v_embed = torch.cat((img_feat, self.l_mask.repeat(B, T, 1)), |
| 96 | dim=1) # v |
| 97 | padding_mask = _get_mask( |
| 98 | self.max_length + zeros_len, |
| 99 | self.max_length) # 对tokens长度以外的padding # B, maxlen maxlen |
| 100 | v_mask = torch.zeros((1, 1, self.max_length, L), |
| 101 | device=img_feat.device).tile([B, 1, 1, |
| 102 | 1]) # maxlen L |
| 103 | mask = torch.cat((v_mask, padding_mask), 3) |
| 104 | v_logits = self.forward_decoder(query, v_embed, mask=mask) |
| 105 | |
| 106 | # 2. language decode |
| 107 | if self.training and self.pretraining: |
| 108 | tgt = torch.where(data[0] == self.ignore_index, 0, data[0]) |
| 109 | tokens = F.one_hot(tgt, num_classes=self.out_channels) |
| 110 | tokens = tokens.float() |
| 111 | lengths = data[-1] |
| 112 | else: |
| 113 | tokens = torch.softmax(v_logits, dim=-1) |
| 114 | lengths = _get_length(v_logits) |
| 115 | tokens = tokens.detach() |
| 116 | token_embed = self.proj(tokens) # (N, T, E) |
| 117 | token_embed = self.token_encoder(token_embed) # (T, N, E) |
| 118 | token_embed = token_embed + self.l_embeding |
| 119 | |
| 120 | padding_mask = _get_mask(lengths, |
| 121 | self.max_length) # 对tokens长度以外的padding |
| 122 | mask = torch.cat((v_mask, padding_mask), 3) |
| 123 | l_embed = torch.cat((self.v_mask.repeat(B, L, 1), token_embed), dim=1) |
| 124 | l_logits = self.forward_decoder(query, l_embed, mask=mask) |
| 125 | |
| 126 | # 3. vision language decode |
| 127 | vl_embed = torch.cat((img_feat, token_embed), dim=1) |
| 128 | vl_logits = self.forward_decoder(query, vl_embed, mask=mask) |
| 129 | |
| 130 | if self.training: |
| 131 | return {'align': [vl_logits], 'lang': l_logits, 'vision': v_logits} |
| 132 | else: |
| 133 | return F.softmax(vl_logits, -1) |
nothing calls this directly
no test coverage detected