Image to Patch Embedding Args: 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. norm_layer (nn.Module, optional): Normalization
| 366 | # return flops |
| 367 | |
| 368 | class PatchEmbed(nn.Module): |
| 369 | """ Image to Patch Embedding |
| 370 | |
| 371 | Args: |
| 372 | patch_size (int): Patch token size. Default: 4. |
| 373 | in_chans (int): Number of input image channels. Default: 3. |
| 374 | embed_dim (int): Number of linear projection output channels. Default: 96. |
| 375 | norm_layer (nn.Module, optional): Normalization layer. Default: None |
| 376 | use_conv_embed (bool): Whether use overlapped convolution for patch embedding. Default: False |
| 377 | is_stem (bool): Is the stem block or not. |
| 378 | """ |
| 379 | |
| 380 | def __init__(self, patch_size=4, in_chans=3, embed_dim=96, norm_layer=None, use_conv_embed=False, is_stem=False, use_pre_norm=False): |
| 381 | super().__init__() |
| 382 | patch_size = to_2tuple(patch_size) |
| 383 | self.patch_size = patch_size |
| 384 | |
| 385 | self.in_chans = in_chans |
| 386 | self.embed_dim = embed_dim |
| 387 | self.use_pre_norm = use_pre_norm |
| 388 | |
| 389 | if use_conv_embed: |
| 390 | # if we choose to use conv embedding, then we treat the stem and non-stem differently |
| 391 | if is_stem: |
| 392 | kernel_size = 7; padding = 3; stride = 4 |
| 393 | else: |
| 394 | kernel_size = 3; padding = 1; stride = 2 |
| 395 | self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=kernel_size, stride=stride, padding=padding) |
| 396 | else: |
| 397 | self.proj = nn.Conv2d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size) |
| 398 | |
| 399 | if self.use_pre_norm: |
| 400 | if norm_layer is not None: |
| 401 | self.norm = norm_layer(in_chans) |
| 402 | else: |
| 403 | self.norm = None |
| 404 | else: |
| 405 | if norm_layer is not None: |
| 406 | self.norm = norm_layer(embed_dim) |
| 407 | else: |
| 408 | self.norm = None |
| 409 | |
| 410 | def forward(self, x): |
| 411 | """Forward function.""" |
| 412 | B, C, H, W = x.size() |
| 413 | if W % self.patch_size[1] != 0: |
| 414 | x = F.pad(x, (0, self.patch_size[1] - W % self.patch_size[1])) |
| 415 | if H % self.patch_size[0] != 0: |
| 416 | x = F.pad(x, (0, 0, 0, self.patch_size[0] - H % self.patch_size[0])) |
| 417 | |
| 418 | if self.use_pre_norm: |
| 419 | if self.norm is not None: |
| 420 | x = x.flatten(2).transpose(1, 2) # B Ph*Pw C |
| 421 | x = self.norm(x).transpose(1, 2).view(B, C, H, W) |
| 422 | x = self.proj(x) |
| 423 | else: |
| 424 | x = self.proj(x) # B C Wh Ww |
| 425 | if self.norm is not None: |