r""" Image to Patch Embedding Args: img_size (int): Image size. Default: 224. patch_size (int): Patch token size. Default: 4. in_chans (int): Number of input image channels. Default: 3. embed_dim (int): Number of linear projection output channels. Default: 96.
| 313 | |
| 314 | |
| 315 | class PatchEmbed(nn.Module): |
| 316 | r""" Image to Patch Embedding |
| 317 | |
| 318 | Args: |
| 319 | img_size (int): Image size. Default: 224. |
| 320 | patch_size (int): Patch token size. Default: 4. |
| 321 | in_chans (int): Number of input image channels. Default: 3. |
| 322 | embed_dim (int): Number of linear projection output channels. Default: 96. |
| 323 | norm_layer (nn.Module, optional): Normalization layer. Default: None |
| 324 | """ |
| 325 | |
| 326 | def __init__(self, img_size=(224, 224), patch_size=4, in_chans=3, embed_dim=96, use_conv_embed=False, |
| 327 | norm_layer=None, is_stem=False): |
| 328 | super().__init__() |
| 329 | patch_size = to_2tuple(patch_size) |
| 330 | patches_resolution = [img_size[0] // patch_size[0], img_size[1] // patch_size[1]] |
| 331 | self.img_size = img_size |
| 332 | self.patch_size = patch_size |
| 333 | self.patches_resolution = patches_resolution |
| 334 | self.num_patches = patches_resolution[0] * patches_resolution[1] |
| 335 | |
| 336 | self.in_chans = in_chans |
| 337 | self.embed_dim = embed_dim |
| 338 | |
| 339 | if use_conv_embed: |
| 340 | # if we choose to use conv embedding, then we treat the stem and non-stem differently |
| 341 | if is_stem: |
| 342 | kernel_size = 7; |
| 343 | padding = 2; |
| 344 | stride = 4 |
| 345 | else: |
| 346 | kernel_size = 3; |
| 347 | padding = 1; |
| 348 | stride = 2 |
| 349 | self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=kernel_size, stride=stride, padding=padding) |
| 350 | else: |
| 351 | self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) |
| 352 | |
| 353 | if norm_layer is not None: |
| 354 | self.norm = norm_layer(embed_dim) |
| 355 | else: |
| 356 | self.norm = None |
| 357 | |
| 358 | def forward(self, x): |
| 359 | B, C, H, W = x.shape |
| 360 | |
| 361 | x = self.proj(x) |
| 362 | H, W = x.shape[2:] |
| 363 | x = x.flatten(2).transpose(1, 2) # B Ph*Pw C |
| 364 | if self.norm is not None: |
| 365 | x = self.norm(x) |
| 366 | return x, H, W |
| 367 | |
| 368 | def flops(self): |
| 369 | Ho, Wo = self.patches_resolution |
| 370 | flops = Ho * Wo * self.embed_dim * self.in_chans * (self.patch_size[0] * self.patch_size[1]) |
| 371 | if self.norm is not None: |
| 372 | flops += Ho * Wo * self.embed_dim |