MCPcopy Create free account
hub / github.com/Monalissaa/DisenDiff / FrozenCLIPEmbedderWrapper

Class FrozenCLIPEmbedderWrapper

src/custom_modules.py:225–312  ·  view source on GitHub ↗

Uses the CLIP transformer encoder for text (from Hugging Face)

Source from the content-addressed store, hash-verified

223
224
225class FrozenCLIPEmbedderWrapper(AbstractEncoder):
226 """Uses the CLIP transformer encoder for text (from Hugging Face)"""
227 def __init__(self, modifier_token, version="openai/clip-vit-large-patch14", device="cuda", max_length=77):
228 super().__init__()
229 self.tokenizer = CLIPTokenizer.from_pretrained(version)
230 self.transformer = CLIPTextModel.from_pretrained(version)
231 self.device = device
232 self.max_length = max_length
233 self.modifier_token = modifier_token
234 if '+' in self.modifier_token:
235 self.modifier_token = self.modifier_token.split('+')
236 else:
237 self.modifier_token = [self.modifier_token]
238
239 self.add_token()
240 self.freeze()
241
242 def add_token(self):
243 self.modifier_token_id = []
244 for each_modifier_token in self.modifier_token:
245 num_added_tokens = self.tokenizer.add_tokens(each_modifier_token)
246 modifier_token_id = self.tokenizer.convert_tokens_to_ids(each_modifier_token)
247 self.modifier_token_id.append(modifier_token_id)
248
249 self.transformer.resize_token_embeddings(len(self.tokenizer))
250 token_embeds = self.transformer.get_input_embeddings().weight.data
251 token_embeds[self.modifier_token_id[-1]] = torch.nn.Parameter(token_embeds[42170], requires_grad=True)
252 if len(self.modifier_token) == 2:
253 token_embeds[self.modifier_token_id[-2]] = torch.nn.Parameter(token_embeds[47629], requires_grad=True)
254 if len(self.modifier_token) == 3:
255 token_embeds[self.modifier_token_id[-3]] = torch.nn.Parameter(token_embeds[43514], requires_grad=True)
256
257 def custom_forward(self, hidden_states, input_ids):
258 r"""
259 Returns:
260 """
261 input_shape = hidden_states.size()
262 bsz, seq_len = input_shape[:2]
263 if version.parse(transformers.__version__) >= version.parse('4.21'):
264 causal_attention_mask = self.transformer.text_model._build_causal_attention_mask(bsz, seq_len, hidden_states.dtype).to(
265 hidden_states.device
266 )
267 else:
268 causal_attention_mask = self.transformer.text_model._build_causal_attention_mask(bsz, seq_len).to(
269 hidden_states.device
270 )
271
272 encoder_outputs = self.transformer.text_model.encoder(
273 inputs_embeds=hidden_states,
274 causal_attention_mask=causal_attention_mask,
275 )
276
277 last_hidden_state = encoder_outputs[0]
278 last_hidden_state = self.transformer.text_model.final_layer_norm(last_hidden_state)
279
280 return last_hidden_state
281
282 def freeze(self):

Callers 1

custom_modules.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected