MCPcopy Create free account
hub / github.com/OpenMatch/UniVL-DR / encoder_model

Class encoder_model

CLIP-DPR/models.py:10–29  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

8
9
10class 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

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected