| 16 | |
| 17 | |
| 18 | class ResidualConvBlock(nn.Module): |
| 19 | def __init__( |
| 20 | self, |
| 21 | in_channels: int, |
| 22 | out_channels: int = None, |
| 23 | hidden_channels: int = None, |
| 24 | kernel_size: int = 3, |
| 25 | padding_mode: str = 'replicate', |
| 26 | activation: Literal['relu', 'leaky_relu', 'silu', 'elu'] = 'relu', |
| 27 | in_norm: Literal['group_norm', 'layer_norm', 'instance_norm', 'none'] = 'layer_norm', |
| 28 | hidden_norm: Literal['group_norm', 'layer_norm', 'instance_norm'] = 'group_norm', |
| 29 | ): |
| 30 | super(ResidualConvBlock, self).__init__() |
| 31 | if out_channels is None: |
| 32 | out_channels = in_channels |
| 33 | if hidden_channels is None: |
| 34 | hidden_channels = in_channels |
| 35 | |
| 36 | if activation =='relu': |
| 37 | activation_cls = nn.ReLU |
| 38 | elif activation == 'leaky_relu': |
| 39 | activation_cls = functools.partial(nn.LeakyReLU, negative_slope=0.2) |
| 40 | elif activation =='silu': |
| 41 | activation_cls = nn.SiLU |
| 42 | elif activation == 'elu': |
| 43 | activation_cls = nn.ELU |
| 44 | else: |
| 45 | raise ValueError(f'Unsupported activation function: {activation}') |
| 46 | |
| 47 | self.layers = nn.Sequential( |
| 48 | nn.GroupNorm(in_channels // 32, in_channels) if in_norm == 'group_norm' else \ |
| 49 | nn.GroupNorm(1, in_channels) if in_norm == 'layer_norm' else \ |
| 50 | nn.InstanceNorm2d(in_channels) if in_norm == 'instance_norm' else \ |
| 51 | nn.Identity(), |
| 52 | activation_cls(), |
| 53 | nn.Conv2d(in_channels, hidden_channels, kernel_size=kernel_size, padding=kernel_size // 2, padding_mode=padding_mode), |
| 54 | nn.GroupNorm(hidden_channels // 32, hidden_channels) if hidden_norm == 'group_norm' else \ |
| 55 | nn.GroupNorm(1, hidden_channels) if hidden_norm == 'layer_norm' else \ |
| 56 | nn.InstanceNorm2d(hidden_channels) if hidden_norm == 'instance_norm' else\ |
| 57 | nn.Identity(), |
| 58 | activation_cls(), |
| 59 | nn.Conv2d(hidden_channels, out_channels, kernel_size=kernel_size, padding=kernel_size // 2, padding_mode=padding_mode) |
| 60 | ) |
| 61 | |
| 62 | self.skip_connection = nn.Conv2d(in_channels, out_channels, kernel_size=1, padding=0) if in_channels != out_channels else nn.Identity() |
| 63 | |
| 64 | def forward(self, x): |
| 65 | skip = self.skip_connection(x) |
| 66 | x = self.layers(x) |
| 67 | x = x + skip |
| 68 | return x |
| 69 | |
| 70 | |
| 71 | class DINOv2Encoder(nn.Module): |