| 26 | |
| 27 | |
| 28 | class PatchEmbed(torch.nn.Module): |
| 29 | def __init__(self, patch_size=2, in_channels=16, embed_dim=1536, pos_embed_max_size=192): |
| 30 | super().__init__() |
| 31 | self.pos_embed_max_size = pos_embed_max_size |
| 32 | self.patch_size = patch_size |
| 33 | |
| 34 | self.proj = torch.nn.Conv2d(in_channels, embed_dim, kernel_size=(patch_size, patch_size), stride=patch_size) |
| 35 | self.pos_embed = torch.nn.Parameter(torch.zeros(1, self.pos_embed_max_size, self.pos_embed_max_size, embed_dim)) |
| 36 | |
| 37 | def cropped_pos_embed(self, height, width): |
| 38 | height = height // self.patch_size |
| 39 | width = width // self.patch_size |
| 40 | top = (self.pos_embed_max_size - height) // 2 |
| 41 | left = (self.pos_embed_max_size - width) // 2 |
| 42 | spatial_pos_embed = self.pos_embed[:, top : top + height, left : left + width, :].flatten(1, 2) |
| 43 | return spatial_pos_embed |
| 44 | |
| 45 | def forward(self, latent): |
| 46 | height, width = latent.shape[-2:] |
| 47 | latent = self.proj(latent) |
| 48 | latent = latent.flatten(2).transpose(1, 2) |
| 49 | pos_embed = self.cropped_pos_embed(height, width) |
| 50 | return latent + pos_embed |
| 51 | |
| 52 | |
| 53 | |