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

Class FrozenCLIPEmbedder

datasets/clip.py:13–48  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

11
12
13class FrozenCLIPEmbedder(AbstractEncoder):
14 """Uses the CLIP transformer encoder for text (from Hugging Face)"""
15
16 def __init__(
17 self, version="openai/clip-vit-large-patch14", device="cuda", max_length=77
18 ):
19 super().__init__()
20 self.tokenizer = CLIPTokenizer.from_pretrained(version)
21 self.transformer = CLIPTextModel.from_pretrained(version)
22 self.device = device
23 self.max_length = max_length
24 self.freeze()
25
26 def freeze(self):
27 self.transformer = self.transformer.eval()
28 for param in self.parameters():
29 param.requires_grad = False
30
31 def forward(self, text):
32 batch_encoding = self.tokenizer(
33 text,
34 truncation=True,
35 max_length=self.max_length,
36 return_length=True,
37 return_overflowing_tokens=False,
38 padding="max_length",
39 return_tensors="pt",
40 )
41 tokens = batch_encoding["input_ids"].to(self.device)
42 outputs = self.transformer(input_ids=tokens)
43
44 z = outputs.last_hidden_state
45 return z
46
47 def encode(self, text):
48 return self(text)

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected