| 25 | __all__ = ['filip_res50', 'filip_vitb32'] |
| 26 | |
| 27 | class FILIP(CLIP): |
| 28 | def __init__(self,image_encode, text_encode, use_allgather, nn_size=2**16, nn_topk=1, \ |
| 29 | return_dense=False, return_caption=False, return_nn_bank=False, text_mask_type=None, \ |
| 30 | EDA=True, feature_dim=1024, embed_dim=768, forward_type='split', dense_mapping_image=2048, \ |
| 31 | dense_mapping_language=512, dense_embed_dim=256, mask_rate=0.75, patch_number=14, \ |
| 32 | text_mae_feature=False, return_simsiam=False, two_view=False, sparse=False, select_topk=False): |
| 33 | super(FILIP, self).__init__(image_encode, text_encode, use_allgather) |
| 34 | self.return_dense = return_dense |
| 35 | self.return_caption = return_caption |
| 36 | |
| 37 | self.text_mask_type = text_mask_type |
| 38 | self.select_topk = select_topk |
| 39 | if self.return_dense: |
| 40 | self.image_mapping = nn.Linear(dense_mapping_image, dense_embed_dim) |
| 41 | self.text_mapping = nn.Linear(dense_mapping_language, dense_embed_dim) |
| 42 | |
| 43 | self.logit_scale_dense = nn.Parameter(torch.ones([])) |
| 44 | nn.init.constant_(self.logit_scale_dense, np.log(1/0.07)) |
| 45 | |
| 46 | if self.encode_text.text_encode_type == 'Transformer': |
| 47 | self.sos_index = self.encode_text.tokenizer.encoder["<|startoftext|>"] |
| 48 | self.padding_idx = 0 |
| 49 | if self.return_caption: |
| 50 | self.caption_module = TransformerDecoderTextualHead(visual_feature_size=2048, vocab_size=self.encode_text.vocab_size, padding_idx=self.padding_idx) |
| 51 | else: |
| 52 | self.caption_module = None |
| 53 | if text_mask_type is not None: |
| 54 | enc_dim = self.encode_text.text_projection.weight.shape[-1] |
| 55 | self.text_label_predictor = nn.Linear(enc_dim, self.encode_text.vocab_size) |
| 56 | |
| 57 | def encode_text_dense(self, texts, return_dense=True): |
| 58 | text_features, word_features = self.encode_text(texts, return_dense=return_dense) |
| 59 | word_features_d = self.text_mapping(word_features) |
| 60 | return word_features_d |
| 61 | |
| 62 | def encode_image_dense(self, image): |
| 63 | image_features, image_features_dense = self.visual(image.type(self.dtype), return_dense=True) |
| 64 | image_features_dense = self.image_mapping(image_features_dense) |
| 65 | return image_features_dense |
| 66 | |
| 67 | def encode_image(self, image, return_all=False): |
| 68 | output = self.visual(image.type(self.dtype), return_dense=return_all) |
| 69 | return output |
| 70 | |
| 71 | def get_weighted_dense_logits(self, dense_feat_1, dense_feat_2, top_k=16): |
| 72 | dense_feat_1 = dense_feat_1 / dense_feat_1.norm(dim=-1, keepdim=True) |
| 73 | dense_feat_2 = dense_feat_2 / dense_feat_2.norm(dim=-1, keepdim=True) |
| 74 | |
| 75 | logit_scale_dense = self.logit_scale_dense.exp() |
| 76 | |
| 77 | |
| 78 | if self.select_topk: |
| 79 | dense_feat_cross_logit = torch.matmul(dense_feat_1, dense_feat_2.permute(0, 2, 1)) |
| 80 | _, dense_id_1 = torch.topk(dense_feat_cross_logit.sum(dim=2), dim=1, k=top_k) |
| 81 | _, dense_id_2 = torch.topk(dense_feat_cross_logit.sum(dim=1), dim=1, k=top_k) |
| 82 | bs, n1 = dense_feat_1.shape[:2] |
| 83 | dense_id_1 = dense_id_1 + (torch.arange(bs) * n1).to(dense_id_1.device)[:, None] |
| 84 | selected_feat_1 = dense_feat_1.reshape(bs * n1, -1)[dense_id_1].reshape(bs, top_k, -1) |
no outgoing calls
no test coverage detected