| 35 | positives: Tuple[str] = ("",) |
| 36 | |
| 37 | class 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 | |