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

Class PatchEmbed

diffsynth/models/sd3_dit.py:28–50  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

26
27
28class 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

Callers 2

__init__Method · 0.70
__init__Method · 0.50

Calls

no outgoing calls

Tested by

no test coverage detected