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

Class CLIP

prototype/model/slip.py:112–206  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

110
111
112class CLIP(nn.Module):
113 def __init__(self,image_encode, text_encode, use_allgather):
114 super().__init__()
115 self.use_allgather = use_allgather
116 self.visual =image_encode
117 self.text_encoder = text_encode
118 self.logit_scale = nn.Parameter(torch.ones([1]))
119 # self.logit_scale = nn.Parameter(torch.ones([]))
120 nn.init.constant_(self.logit_scale, np.log(1/0.07))
121 #nn.init.normal_(self.text_projection, std=self.transformer.width ** -0.5)
122
123 def text_parameters(self):
124 param = [self.logit_scale]
125 if self.text_encoder.text_encode_type == 'Transformer':
126 param.append(self.text_encoder.positional_embedding)
127 elif self.text_encoder.text_encode_type == 'Bert':
128 # print('Bert', self.text_encoder.text_transformer.cls.predictions, flush=True)
129 # param.extend([self.text_encoder.text_transformer.cls.predictions.decoder.weight,
130 # self.text_encoder.text_transformer.cls.predictions.bias])
131 param.extend([self.text_encoder.text_transformer.cls.predictions.bias])
132 return param
133
134 def text_modules(self):
135 if self.text_encoder.text_encode_type == 'Transformer':
136 return [self.text_encoder.transformer, self.text_encoder.text_projection, self.text_encoder.token_embedding, self.text_encoder.ln_final]
137 elif self.text_encoder.text_encode_type == 'Bert':
138 # print('Bert', self.text_encoder.text_transformer, flush=True)
139 return [self.text_encoder.text_transformer.bert, self.text_encoder.text_projection,
140 self.text_encoder.text_transformer.cls.predictions.transform]
141 # self.text_encoder.text_transformer.cls.predictions.decoder, # decoder: bias
142 else:
143 import ipdb
144 ipdb.set_trace()
145 return [self.text_encoder.text_transformer, self.text_encoder.text_projection]
146
147 def visual_parameters(self):
148 return []
149
150 def visual_modules(self):
151 return [self.visual]
152
153 @property
154 def dtype(self):
155 try:
156 return self.visual.conv1.weight.dtype
157 except:
158 try:
159 return self.visual.head.weight.dtype
160 except:
161 try:
162 return self.visual.stem[0].weight.dtype
163 except:
164 return self.text_encoder.text_projection.weight.dtype
165
166 def encode_image(self, image):
167 return self.visual(image.type(self.dtype))
168
169 def sample_captions(self, texts):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected