MCPcopy Create free account
hub / github.com/CompVis/diff2flow / FrozenCLIPEmbedder

Class FrozenCLIPEmbedder

diff2flow/conditioning/encoders.py:224–263  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

222 return self.uncond.to(device)
223
224class 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
266if __name__ == "__main__":

Callers 1

encoders.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected