| 57 | |
| 58 | class DoubleConv(nn.Module): |
| 59 | def __init__(self, in_channels, out_channels, mid_channels=None, residual=False): |
| 60 | super().__init__() |
| 61 | self.residual = residual |
| 62 | if not mid_channels: |
| 63 | mid_channels = out_channels |
| 64 | self.double_conv = nn.Sequential( |
| 65 | nn.Conv2d(in_channels, mid_channels, kernel_size=3, padding=1, bias=False), |
| 66 | nn.GroupNorm(1, mid_channels), |
| 67 | nn.GELU(), |
| 68 | nn.Conv2d(mid_channels, out_channels, kernel_size=3, padding=1, bias=False), |
| 69 | nn.GroupNorm(1, out_channels), |
| 70 | ) |
| 71 | |
| 72 | def forward(self, x): |
| 73 | if self.residual: |