| 611 | |
| 612 | # Input Projection |
| 613 | class InputProj(nn.Module): |
| 614 | def __init__(self, in_channel=3, out_channel=64, kernel_size=3, stride=1, norm_layer=None,act_layer=nn.LeakyReLU): |
| 615 | super().__init__() |
| 616 | self.proj = nn.Sequential( |
| 617 | nn.Conv2d(in_channel, out_channel, kernel_size=3, stride=stride, padding=kernel_size//2), |
| 618 | act_layer(inplace=True) |
| 619 | ) |
| 620 | if norm_layer is not None: |
| 621 | self.norm = norm_layer(out_channel) |
| 622 | else: |
| 623 | self.norm = None |
| 624 | self.in_channel = in_channel |
| 625 | self.out_channel = out_channel |
| 626 | |
| 627 | def forward(self, x): |
| 628 | B, C, H, W = x.shape |
| 629 | x = self.proj(x).flatten(2).transpose(1, 2).contiguous() # B H*W C |
| 630 | if self.norm is not None: |
| 631 | x = self.norm(x) |
| 632 | return x |
| 633 | |
| 634 | def flops(self, H, W): |
| 635 | flops = 0 |
| 636 | # conv |
| 637 | flops += H*W*self.in_channel*self.out_channel*3*3 |
| 638 | |
| 639 | if self.norm is not None: |
| 640 | flops += H*W*self.out_channel |
| 641 | print("Input_proj:{%.2f}"%(flops/1e9)) |
| 642 | return flops |
| 643 | |
| 644 | # Output Projection |
| 645 | class OutputProj(nn.Module): |