| 3 | |
| 4 | |
| 5 | class CLIPVisionEmbeddings(torch.nn.Module): |
| 6 | def __init__(self, embed_dim=1280, image_size=224, patch_size=14, num_channels=3): |
| 7 | super().__init__() |
| 8 | |
| 9 | # class_embeds (This is a fixed tensor) |
| 10 | self.class_embedding = torch.nn.Parameter(torch.randn(1, 1, embed_dim)) |
| 11 | |
| 12 | # position_embeds |
| 13 | self.patch_embedding = torch.nn.Conv2d(in_channels=num_channels, out_channels=embed_dim, kernel_size=patch_size, stride=patch_size, bias=False) |
| 14 | |
| 15 | # position_embeds (This is a fixed tensor) |
| 16 | self.position_embeds = torch.nn.Parameter(torch.zeros(1, (image_size // patch_size) ** 2 + 1, embed_dim)) |
| 17 | |
| 18 | def forward(self, pixel_values): |
| 19 | batch_size = pixel_values.shape[0] |
| 20 | patch_embeds = self.patch_embedding(pixel_values) |
| 21 | patch_embeds = patch_embeds.flatten(2).transpose(1, 2) |
| 22 | class_embeds = self.class_embedding.repeat(batch_size, 1, 1) |
| 23 | embeddings = torch.cat([class_embeds, patch_embeds], dim=1) + self.position_embeds |
| 24 | return embeddings |
| 25 | |
| 26 | |
| 27 | class SVDImageEncoder(torch.nn.Module): |