Uses the CLIP transformer encoder for text (from Hugging Face)
| 224 | |
| 225 | |
| 226 | class 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 |