MCPcopy Create free account
hub / github.com/42dot/VFDepth / DepthDecoder

Class DepthDecoder

network/fusion_depthnet.py:97–145  ·  view source on GitHub ↗

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))

Source from the content-addressed store, hash-verified

95
96
97class 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

Callers 2

__init__Method · 0.85
__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected