This class decodes encoded 2D features to estimate depth map. Unlike monodepth depth decoder, we decode features with corresponding level we used to project features in 3D (default: level 2(H/4, W/4))
| 95 | |
| 96 | |
| 97 | class DepthDecoder(nn.Module): |
| 98 | """ |
| 99 | This class decodes encoded 2D features to estimate depth map. |
| 100 | Unlike monodepth depth decoder, we decode features with corresponding level we used to project features in 3D (default: level 2(H/4, W/4)) |
| 101 | """ |
| 102 | def __init__(self, level_in, num_ch_enc, num_ch_dec, scales=range(2), use_skips=False): |
| 103 | super(DepthDecoder, self).__init__() |
| 104 | |
| 105 | self.num_output_channels = 1 |
| 106 | self.scales = scales |
| 107 | self.use_skips = use_skips |
| 108 | |
| 109 | self.level_in = level_in |
| 110 | self.num_ch_enc = num_ch_enc |
| 111 | self.num_ch_dec = num_ch_dec |
| 112 | |
| 113 | self.convs = OrderedDict() |
| 114 | for i in range(self.level_in, -1, -1): |
| 115 | num_ch_in = self.num_ch_enc[-1] if i == self.level_in else self.num_ch_dec[i + 1] |
| 116 | num_ch_out = self.num_ch_dec[i] |
| 117 | self.convs[('upconv', i, 0)] = conv2d(num_ch_in, num_ch_out, kernel_size=3, nonlin = 'ELU') |
| 118 | |
| 119 | num_ch_in = self.num_ch_dec[i] |
| 120 | if self.use_skips and i > 0: |
| 121 | num_ch_in += self.num_ch_enc[i - 1] |
| 122 | num_ch_out = self.num_ch_dec[i] |
| 123 | self.convs[('upconv', i, 1)] = conv2d(num_ch_in, num_ch_out, kernel_size=3, nonlin = 'ELU') |
| 124 | |
| 125 | for s in self.scales: |
| 126 | self.convs[('dispconv', s)] = conv2d(self.num_ch_dec[s], self.num_output_channels, 3, nonlin = None) |
| 127 | |
| 128 | self.decoder = nn.ModuleList(list(self.convs.values())) |
| 129 | self.sigmoid = nn.Sigmoid() |
| 130 | |
| 131 | def forward(self, input_features): |
| 132 | outputs = {} |
| 133 | |
| 134 | # decode |
| 135 | x = input_features[-1] |
| 136 | for i in range(self.level_in, -1, -1): |
| 137 | x = self.convs[('upconv', i, 0)](x) |
| 138 | x = [upsample(x)] |
| 139 | if self.use_skips and i > 0: |
| 140 | x += [input_features[i - 1]] |
| 141 | x = torch.cat(x, 1) |
| 142 | x = self.convs[('upconv', i, 1)](x) |
| 143 | if i in self.scales: |
| 144 | outputs[('disp', i)] = self.sigmoid(self.convs[('dispconv', i)](x)) |
| 145 | return outputs |