MCPcopy Create free account
hub / github.com/adobe-research/custom-diffusion / FrozenCLIPEmbedderWrapper

Class FrozenCLIPEmbedderWrapper

src/custom_modules.py:226–315  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

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

Callers 1

custom_modules.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected