| 8 | |
| 9 | |
| 10 | class encoder_model(nn.Module): |
| 11 | def __init__(self, encoder): |
| 12 | super(encoder_model, self).__init__() |
| 13 | self.encoder = encoder |
| 14 | self.project_layer = nn.Linear(1024, 512) |
| 15 | self.logit_scale = self.encoder.logit_scale |
| 16 | |
| 17 | def forward(self, text=None, image=None, captions=None): |
| 18 | if text != None: |
| 19 | embeddings = self.encoder(text=text, image=None) |
| 20 | elif image != None: |
| 21 | embeddings = self.encoder(image=image, text=None) |
| 22 | if captions != None: |
| 23 | cap_embeddings = self.encoder(text=captions, image=None) |
| 24 | embeddings = torch.cat([embeddings, cap_embeddings], -1) |
| 25 | embeddings = self.project_layer(embeddings) |
| 26 | else: |
| 27 | raise ("error!") |
| 28 | embeddings = F.normalize(embeddings, dim=-1) |
| 29 | return embeddings |
nothing calls this directly
no outgoing calls
no test coverage detected