MCPcopy Create free account
hub / github.com/OpenImagingLab/FlashVSR / CLIPVisionEmbeddings

Class CLIPVisionEmbeddings

diffsynth/models/svd_image_encoder.py:5–24  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

3
4
5class 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
27class SVDImageEncoder(torch.nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected