MCPcopy Create free account
hub / github.com/Sense-GVT/DeCLIP / FILIP

Class FILIP

prototype/model/filip.py:27–142  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

25__all__ = ['filip_res50', 'filip_vitb32']
26
27class 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)

Callers 2

filip_res50Function · 0.85
filip_vitb32Function · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected