(self, text_list)
| 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 | |
| 95 | def get_relevancy(self, embed: torch.Tensor, positive_id: int) -> torch.Tensor: |
| 96 | phrases_embeds = torch.cat([self.pos_embeds, self.neg_embeds], dim=0) |
no test coverage detected