MCPcopy Create free account
hub / github.com/Atrovast/THGS / OpenCLIPNetwork

Class OpenCLIPNetwork

scripts/image_encoding.py:37–110  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

35 positives: Tuple[str] = ("",)
36
37class OpenCLIPNetwork(nn.Module):
38 def __init__(self, config: OpenCLIPNetworkConfig):
39 super().__init__()
40 self.config = config
41 self.process = torchvision.transforms.Compose(
42 [
43 torchvision.transforms.Resize((224, 224)),
44 torchvision.transforms.Normalize(
45 mean=[0.48145466, 0.4578275, 0.40821073],
46 std=[0.26862954, 0.26130258, 0.27577711],
47 ),
48 ]
49 )
50 model, _, _ = open_clip.create_model_and_transforms(
51 self.config.clip_model_type, # e.g., ViT-B-16
52 pretrained=self.config.clip_model_pretrained, # e.g., laion2b_s34b_b88k
53 precision="fp16",
54 )
55 model.eval()
56 self.tokenizer = open_clip.get_tokenizer(self.config.clip_model_type)
57 self.model = model.to("cuda")
58 self.clip_n_dims = self.config.clip_n_dims
59
60 self.positives = self.config.positives
61 self.negatives = self.config.negatives
62 with torch.no_grad():
63 tok_phrases = torch.cat([self.tokenizer(phrase) for phrase in self.positives]).to("cuda")
64 self.pos_embeds = model.encode_text(tok_phrases)
65 tok_phrases = torch.cat([self.tokenizer(phrase) for phrase in self.negatives]).to("cuda")
66 self.neg_embeds = model.encode_text(tok_phrases)
67 self.pos_embeds /= self.pos_embeds.norm(dim=-1, keepdim=True)
68 self.neg_embeds /= self.neg_embeds.norm(dim=-1, keepdim=True)
69
70 assert (
71 self.pos_embeds.shape[1] == self.neg_embeds.shape[1]
72 ), "Positive and negative embeddings must have the same dimensionality"
73 assert (
74 self.pos_embeds.shape[1] == self.clip_n_dims
75 ), "Embedding dimensionality must match the model dimensionality"
76
77 @property
78 def name(self) -> str:
79 return "openclip_{}_{}".format(self.config.clip_model_type, self.config.clip_model_pretrained)
80
81 @property
82 def embedding_dim(self) -> int:
83 return self.config.clip_n_dims
84
85 def gui_cb(self,element):
86 self.set_positives(element.value.split(";"))
87
88 def set_positives(self, text_list):
89 self.positives = text_list
90 with torch.no_grad():
91 tok_phrases = torch.cat([self.tokenizer(phrase) for phrase in self.positives]).to("cuda")
92 self.pos_embeds = self.model.encode_text(tok_phrases)
93 self.pos_embeds /= self.pos_embeds.norm(dim=-1, keepdim=True)
94

Callers 1

image_encoding.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected