Decoder for multi-resolution encodings.
| 20 | |
| 21 | |
| 22 | class MultiresConvDecoder(BaseDecoder): |
| 23 | """Decoder for multi-resolution encodings.""" |
| 24 | |
| 25 | def __init__( |
| 26 | self, |
| 27 | dims_encoder: Iterable[int], |
| 28 | dims_decoder: Iterable[int] | int, |
| 29 | grad_checkpointing: bool = False, |
| 30 | upsampling_mode: UpsamplingMode = "transposed_conv", |
| 31 | ): |
| 32 | """Initialize multiresolution convolutional decoder. |
| 33 | |
| 34 | Args: |
| 35 | dims_encoder: Expected dims at each level from the encoder. |
| 36 | dims_decoder: Dim of decoder features. |
| 37 | grad_checkpointing: Whether to checkpoint gradient during training. |
| 38 | upsampling_mode: What method to use for upsampling. |
| 39 | """ |
| 40 | super().__init__() |
| 41 | self.dims_encoder = list(dims_encoder) |
| 42 | |
| 43 | if isinstance(dims_decoder, int): |
| 44 | self.dims_decoder = [dims_decoder] * len(self.dims_encoder) |
| 45 | else: |
| 46 | self.dims_decoder = list(dims_decoder) |
| 47 | |
| 48 | if len(self.dims_decoder) != len(self.dims_encoder): |
| 49 | raise ValueError("Received dims_encoder and dims_decoder of different sizes.") |
| 50 | |
| 51 | self.dim_out = self.dims_decoder[0] |
| 52 | |
| 53 | num_encoders = len(self.dims_encoder) |
| 54 | |
| 55 | # At the highest resolution, i.e. level 0, we apply projection w/ 1x1 convolution |
| 56 | # when the dimensions mismatch. Otherwise we do not do anything, which is |
| 57 | # the default behavior of monodepth. |
| 58 | conv0 = ( |
| 59 | nn.Conv2d(self.dims_encoder[0], self.dims_decoder[0], kernel_size=1, bias=False) |
| 60 | if self.dims_encoder[0] != self.dims_decoder[0] |
| 61 | else nn.Identity() |
| 62 | ) |
| 63 | |
| 64 | convs = [conv0] |
| 65 | for i in range(1, num_encoders): |
| 66 | convs.append( |
| 67 | nn.Conv2d( |
| 68 | self.dims_encoder[i], |
| 69 | self.dims_decoder[i], |
| 70 | kernel_size=3, |
| 71 | stride=1, |
| 72 | padding=1, |
| 73 | bias=False, |
| 74 | ) |
| 75 | ) |
| 76 | self.convs = nn.ModuleList(convs) |
| 77 | |
| 78 | fusions = [] |
| 79 | for i in range(num_encoders): |
no outgoing calls
no test coverage detected