MCPcopy Create free account
hub / github.com/MiniMax-AI/VTP / PatchEmbed

Class PatchEmbed

vtp/models/layers/embeddings.py:18–83  ·  view source on GitHub ↗

2D image to patch embedding: (B,C,H,W) -> (B,N,D) Args: img_size: Image size. patch_size: Patch token size. in_chans: Number of input image channels. embed_dim: Number of linear projection output channels. norm_layer: Normalization layer.

Source from the content-addressed store, hash-verified

16
17
18class PatchEmbed(nn.Module):
19 """
20 2D image to patch embedding: (B,C,H,W) -> (B,N,D)
21
22 Args:
23 img_size: Image size.
24 patch_size: Patch token size.
25 in_chans: Number of input image channels.
26 embed_dim: Number of linear projection output channels.
27 norm_layer: Normalization layer.
28 """
29
30 def __init__(
31 self,
32 img_size: Union[int, Tuple[int, int]] = 224,
33 patch_size: Union[int, Tuple[int, int]] = 16,
34 in_chans: int = 3,
35 embed_dim: int = 768,
36 norm_layer: Optional[Callable] = None,
37 flatten_embedding: bool = True,
38 ) -> None:
39 super().__init__()
40
41 image_HW = make_2tuple(img_size)
42 patch_HW = make_2tuple(patch_size)
43 patch_grid_size = (
44 image_HW[0] // patch_HW[0],
45 image_HW[1] // patch_HW[1],
46 )
47
48 self.img_size = image_HW
49 self.patch_size = patch_HW
50 self.patches_resolution = patch_grid_size
51 self.num_patches = patch_grid_size[0] * patch_grid_size[1]
52
53 self.in_chans = in_chans
54 self.embed_dim = embed_dim
55
56 self.flatten_embedding = flatten_embedding
57
58 self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_HW, stride=patch_HW)
59 self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
60
61 def forward(self, x: Tensor) -> Tensor:
62 _, _, H, W = x.shape
63
64 x = self.proj(x) # B C H W
65 H, W = x.size(2), x.size(3)
66 x = x.flatten(2).transpose(1, 2) # B HW C
67 x = self.norm(x)
68 if not self.flatten_embedding:
69 x = x.reshape(-1, H, W, self.embed_dim) # B H W C
70 return x
71
72 def flops(self) -> float:
73 Ho, Wo = self.patches_resolution
74 flops = Ho * Wo * self.embed_dim * self.in_chans * (self.patch_size[0] * self.patch_size[1])
75 if self.norm is not None:

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected