(self, text)
| 24 | return |
| 25 | |
| 26 | def encode_text(self, text): |
| 27 | text = self.tokenizer([text] + self.canon).to(self.device) |
| 28 | with torch.no_grad(): |
| 29 | text_features = self.clip_pretrained.encode_text(text).type(torch.float32) |
| 30 | text_features = (text_features / text_features.norm(dim=-1, keepdim=True)).to(self.device) |
| 31 | self.text_feature = text_features |
| 32 | # return text_features |
| 33 | |
| 34 | def compute_similarity(self, semantic_feature): |
| 35 | logit = semantic_feature @ self.text_feature.T |