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 l
| 447 | |
| 448 | |
| 449 | class PatchEmbed(nn.Module): |
| 450 | """ Image to Patch Embedding |
| 451 | Args: |
| 452 | patch_size (int): Patch token size. Default: 4. |
| 453 | in_chans (int): Number of input image channels. Default: 3. |
| 454 | embed_dim (int): Number of linear projection output channels. Default: 96. |
| 455 | norm_layer (nn.Module, optional): Normalization layer. Default: None |
| 456 | """ |
| 457 | def __init__(self, |
| 458 | patch_size=4, |
| 459 | in_chans=3, |
| 460 | embed_dim=96, |
| 461 | norm_layer=None): |
| 462 | super().__init__() |
| 463 | patch_size = to_2tuple(patch_size) |
| 464 | self.patch_size = patch_size |
| 465 | |
| 466 | self.in_chans = in_chans |
| 467 | self.embed_dim = embed_dim |
| 468 | |
| 469 | self.proj = nn.Conv2d(in_chans, |
| 470 | embed_dim, |
| 471 | kernel_size=patch_size, |
| 472 | stride=patch_size) |
| 473 | if norm_layer is not None: |
| 474 | self.norm = norm_layer(embed_dim) |
| 475 | else: |
| 476 | self.norm = None |
| 477 | |
| 478 | def forward(self, x): |
| 479 | """Forward function.""" |
| 480 | # padding |
| 481 | _, _, H, W = x.size() |
| 482 | if W % self.patch_size[1] != 0: |
| 483 | x = F.pad(x, (0, self.patch_size[1] - W % self.patch_size[1])) |
| 484 | if H % self.patch_size[0] != 0: |
| 485 | x = F.pad(x, |
| 486 | (0, 0, 0, self.patch_size[0] - H % self.patch_size[0])) |
| 487 | |
| 488 | x = self.proj(x) # B C Wh Ww |
| 489 | if self.norm is not None: |
| 490 | Wh, Ww = x.size(2), x.size(3) |
| 491 | x = x.flatten(2).transpose(1, 2) |
| 492 | x = self.norm(x) |
| 493 | x = x.transpose(1, 2).view(-1, self.embed_dim, Wh, Ww) |
| 494 | |
| 495 | return x |
| 496 | |
| 497 | |
| 498 | class SwinTransformer(nn.Module): |