(self, texts, norm=True)
| 107 | |
| 108 | # @torch.no_grad() |
| 109 | def forward_language(self, texts, norm=True): |
| 110 | x = self.lang_encoder(*texts) |
| 111 | x = x['last_hidden_state'] |
| 112 | |
| 113 | if self.tokenizer_type == 'clip': |
| 114 | x = x[torch.arange(x.size(0)), texts[0].argmax(dim=-1)] |
| 115 | else: |
| 116 | x = x[:, 0] |
| 117 | |
| 118 | x = x @ self.lang_proj |
| 119 | if norm: |
| 120 | x = x / (x.norm(dim=-1, keepdim=True) + 1e-7) |
| 121 | return x |
| 122 | |
| 123 | def compute_similarity(self, v_emb, name='default'): |
| 124 | v_emb = v_emb / (v_emb.norm(dim=-1, keepdim=True) + 1e-7) |
no outgoing calls
no test coverage detected