Uses the CLIP transformer encoder for text (from huggingface)
| 86 | |
| 87 | |
| 88 | class FrozenCLIPEmbedder(AbstractEncoder): |
| 89 | """Uses the CLIP transformer encoder for text (from huggingface)""" |
| 90 | LAYERS = [ |
| 91 | "last", |
| 92 | "pooled", |
| 93 | "hidden" |
| 94 | ] |
| 95 | def __init__(self, version="openai/clip-vit-large-patch14", device="cuda", max_length=77, |
| 96 | freeze=True, layer="last", layer_idx=None): # clip-vit-base-patch32 |
| 97 | super().__init__() |
| 98 | assert layer in self.LAYERS |
| 99 | self.tokenizer = CLIPTokenizer.from_pretrained(version) |
| 100 | self.transformer = CLIPTextModel.from_pretrained(version) |
| 101 | self.device = device |
| 102 | self.max_length = max_length |
| 103 | if freeze: |
| 104 | self.freeze() |
| 105 | self.layer = layer |
| 106 | self.layer_idx = layer_idx |
| 107 | if layer == "hidden": |
| 108 | assert layer_idx is not None |
| 109 | assert 0 <= abs(layer_idx) <= 12 |
| 110 | |
| 111 | def freeze(self): |
| 112 | self.transformer = self.transformer.eval() |
| 113 | #self.train = disabled_train |
| 114 | for param in self.parameters(): |
| 115 | param.requires_grad = False |
| 116 | |
| 117 | def forward(self, text): |
| 118 | batch_encoding = self.tokenizer(text, truncation=True, max_length=self.max_length, return_length=True, |
| 119 | return_overflowing_tokens=False, padding="max_length", return_tensors="pt") |
| 120 | tokens = batch_encoding["input_ids"].to(self.device) |
| 121 | outputs = self.transformer(input_ids=tokens, output_hidden_states=self.layer=="hidden") |
| 122 | if self.layer == "last": |
| 123 | z = outputs.last_hidden_state |
| 124 | elif self.layer == "pooled": |
| 125 | z = outputs.pooler_output[:, None, :] |
| 126 | else: |
| 127 | z = outputs.hidden_states[self.layer_idx] |
| 128 | return z |
| 129 | |
| 130 | def encode(self, text): |
| 131 | return self(text) |
| 132 | |
| 133 | |
| 134 | class FrozenOpenCLIPEmbedder(AbstractEncoder): |