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

Class DECLIP

prototype/model/declip.py:132–336  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

130 return x
131
132class DECLIP(CLIP):
133 def __init__(self,image_encode, text_encode, use_allgather, nn_size=2**16, nn_topk=1, \
134 return_dense=False, return_simsiam_text=False, return_simsiam_nn_text=False, return_caption=False, return_nn_bank=False, text_mask_type=None,
135 EDA=True, feature_dim=1024, forward_type='split'):
136 super(DECLIP, self).__init__(image_encode, text_encode, use_allgather)
137 # TODO change for r50 checkpoint
138 self.projector = projection_MLP(feature_dim)
139 # self.projector = projection_MLP(1024)
140 self.predictor = prediction_MLP(1024)
141 self.return_dense = return_dense
142 self.return_simsiam_nn_text = return_simsiam_nn_text
143 self.return_nn_bank = return_nn_bank
144 self.return_caption = return_caption
145 self.return_simsiam_text = return_simsiam_text
146 self.return_simsiam_nn_text = return_simsiam_nn_text
147 self.text_mask_type = text_mask_type
148 self.EDA = EDA
149 self.forward_type = forward_type
150 #import gensim
151 #from textaugment import Word2vec
152 #model = gensim.models.KeyedVectors.load_word2vec_format('/mnt/cache/liyangguang/GoogleNews-vectors-negative300.bin.gz', binary=True)
153 #self.word2vec = Word2vec(model=model)
154 from textaugment import EDA
155 self.emd = EDA()
156
157 if self.return_dense:
158 raise NotImplementedError('These are bugs in the model, Please Check The Codes!')
159 self.projector_d = projection_MLP(2048) # dense
160 self.predictor_d = prediction_MLP(1024)
161 if self.return_simsiam_text:
162 self.projector_text = projection_MLP(feature_dim)
163 self.predictor_text = prediction_MLP(1024)
164 if self.return_simsiam_nn_text:
165 self.projector_nn_text = projection_MLP(feature_dim)
166 self.predictor_nn_text = prediction_MLP(1024)
167 if self.return_caption:
168 raise NotImplementedError('Not Available')
169 if text_mask_type is not None:
170 enc_dim = self.encode_text.text_projection.weight.shape[-1]
171 self.text_label_predictor = nn.Linear(enc_dim, self.encode_text.vocab_size)
172 if self.return_nn_bank:
173 #nn.init.normal_(self.text_projection, std=self.transformer.width ** -0.5)
174 self.nn_replacer_img = NNMemoryBankModule(size=nn_size, topk=nn_topk)
175 self.nn_replacer_text = NNMemoryBankModule(size=nn_size, topk=nn_topk)
176
177 def text_modules(self):
178 ret = super(self).text_modules()
179 if self.text_mask_type is not None:
180 ret.append(self.text_label_predictor)
181 return ret
182
183 def visual_modules(self):
184 ret = super(self).visual_modules()
185 ret.extend([self.predictor, self.projector])
186 return ret
187
188 def encode_image(self, image, return_dense=False):
189 if return_dense:

Callers 2

declip_res50Function · 0.85
declip_vitb32Function · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected