Uses the CLIP transformer encoder for text (from Hugging Face)
| 222 | return self.uncond.to(device) |
| 223 | |
| 224 | class FrozenCLIPEmbedder(nn.Module): |
| 225 | """Uses the CLIP transformer encoder for text (from Hugging Face)""" |
| 226 | def __init__(self, version="openai/clip-vit-large-patch14", device="cuda", max_length=77): |
| 227 | super().__init__() |
| 228 | self.tokenizer = CLIPTokenizer.from_pretrained(version) |
| 229 | self.transformer = CLIPTextModel.from_pretrained(version) |
| 230 | self.device = device |
| 231 | self.max_length = max_length |
| 232 | self.freeze() |
| 233 | |
| 234 | self.uncond = None |
| 235 | |
| 236 | def freeze(self): |
| 237 | self.transformer = self.transformer.eval() |
| 238 | for param in self.parameters(): |
| 239 | param.requires_grad = False |
| 240 | |
| 241 | def forward(self, text): |
| 242 | dev = next(self.parameters()).device |
| 243 | batch_encoding = self.tokenizer(text, truncation=True, max_length=self.max_length, return_length=True, |
| 244 | return_overflowing_tokens=False, padding="max_length", return_tensors="pt") |
| 245 | tokens = batch_encoding["input_ids"].to(dev) |
| 246 | outputs = self.transformer(input_ids=tokens) |
| 247 | |
| 248 | z = outputs.last_hidden_state |
| 249 | return z |
| 250 | |
| 251 | def encode(self, text): |
| 252 | return self(text) |
| 253 | |
| 254 | @torch.no_grad() |
| 255 | def get_unconditional_conditioning(self, device="cuda"): |
| 256 | """ |
| 257 | Returns: |
| 258 | torch.Tensor: Unconditional conditioning information for text |
| 259 | of shape (1, max_length, d_model), e.g. (1, 77, 1024) |
| 260 | """ |
| 261 | if self.uncond is None: |
| 262 | self.uncond = self.encode("") |
| 263 | return self.uncond.to(device) |
| 264 | |
| 265 | |
| 266 | if __name__ == "__main__": |