Views
No views yet
CLIPModelImageEncoderHead)
TextEncoderHead)
1import torch
2import torch.nn as nn
3import torch.nn.functional as F
4
5class ImageEncoderHead(nn.Module, PyTorchModelHubMixin):
6 def __init__(self, model):
7 super(ImageEncoderHead, self).__init__()
8 self.model = model
9 for param in self.model.parameters():
10 param.requires_grad = False
11 self.seq1 = nn.Sequential(
12 nn.Linear(768, 1000),
13 nn.Dropout(0.3),
14 nn.ReLU(),
15 nn.Linear(1000, 512),
16 nn.LayerNorm(512),
17 )
18
19 def forward(self, pixel_values):
20 outputs = self.model(pixel_values).pooler_output
21 outputs = self.seq1(outputs)
22 return outputs.contiguous()
23
24class TextEncoderHead(nn.Module, PyTorchModelHubMixin):
25 def __init__(self, model):
26 super(TextEncoderHead, self).__init__()
27 self.model = model
28 for param in self.model.parameters():
29 param.requires_grad = False
30 self.seq1 = nn.Sequential(
31 nn.Flatten(),
32 nn.Linear(768 * 128, 2000),
33 nn.Dropout(0.3),
34 nn.ReLU(),
35 nn.Linear(2000, 512),
36 nn.LayerNorm(512),
37 )
38
39 def forward(self, input_ids, attention_mask):
40 outputs = self.model(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
41 outputs = self.seq1(outputs)
42 return outputs.contiguous()
43
44class CLIPModel(nn.Module, PyTorchModelHubMixin):
45 def __init__(self, text_encoder, image_encoder):
46 super(CLIPModel, self).__init__()
47 self.text_encoder = text_encoder
48 self.image_encoder = image_encoder
49
50 def forward(self, image, input_ids, attention_mask):
51 ie = self.image_encoder(image)
52 te = self.text_encoder(input_ids, attention_mask)
53 return ie, te