MCPcopy Create free account
hub / github.com/apple/ARKitScenes / MSGNet

Class MSGNet

depth_upsampling/models/msg/msg.py:10–63  ·  view source on GitHub ↗

Inspired by: Depth Map Super-Resolution by Deep Multi-Scale Guidance http://personal.ie.cuhk.edu.hk/~ccloy/files/eccv_2016_depth.pdf

Source from the content-addressed store, hash-verified

8
9
10class MSGNet(nn.Module):
11 """
12 Inspired by: Depth Map Super-Resolution by Deep Multi-Scale Guidance
13 http://personal.ie.cuhk.edu.hk/~ccloy/files/eccv_2016_depth.pdf
14 """
15 def __init__(self, upsampling_factor):
16 super().__init__()
17 # initialize indexes for layers
18 self.upsampling_factor = upsampling_factor
19 m = int(np.log2(upsampling_factor))
20
21 # RGB-branch
22 self.rgb_encoder1 = nn.Sequential(ConvPReLu(3, 49, 7, stride=1, padding=3),
23 ConvPReLu(49, 32))
24 self.rgb_encoder_blocks = nn.ModuleList()
25 for i in range(m-1):
26 self.rgb_encoder_blocks.append(nn.Sequential(ConvPReLu(32, 32),
27 nn.MaxPool2d(3, 2, padding=1)))
28
29 # D-branch
30 self.depth_decoder1 = nn.Sequential(ConvPReLu(1, 64, 5, stride=1, padding=2),
31 DeconvPReLu(64, 32, 5, stride=2, padding=2))
32 self.depth_decoder_blocks = nn.ModuleList()
33 for i in range(m-1):
34 self.depth_decoder_blocks.append(nn.Sequential(ConvPReLu(64, 32, 5, stride=1, padding=2),
35 ConvPReLu(32, 32, 5, stride=1, padding=2),
36 DeconvPReLu(32, 32, 5, stride=2, padding=2)))
37
38 self.depth_decoder_n = nn.Sequential(ConvPReLu(64, 32, 5, stride=1, padding=2),
39 ConvPReLu(32, 32, 5, stride=1, padding=2),
40 ConvPReLu(32, 32, 5, stride=1, padding=2),
41 ConvPReLu(32, 1, 5, stride=1, padding=2))
42
43 def forward(self, batch):
44 rgb_img = batch[dataset_keys.COLOR_IMG] / 255
45 low_res_depth = batch[dataset_keys.LOW_RES_DEPTH_IMG]
46 min_d = low_res_depth.amin((1, 2, 3), keepdim=True)
47 max_d = low_res_depth.amax((1, 2, 3), keepdim=True)
48 low_res_depth_norm = (low_res_depth - min_d) / ((max_d - min_d) + 1e-8)
49 low_res_upsampled = F.interpolate(low_res_depth_norm, rgb_img.shape[2:], mode='bicubic')
50
51 rgb_features = [self.rgb_encoder1(rgb_img), ]
52 for block in self.rgb_encoder_blocks:
53 rgb_features.append(block(rgb_features[-1]))
54
55 rec = self.depth_decoder1(low_res_depth_norm)
56 for i, block in enumerate(self.depth_decoder_blocks):
57 rec = torch.cat((rec, rgb_features[-(i + 1)]), 1)
58 rec = block(rec)
59 rec = torch.cat((rec, rgb_features[0]), 1)
60 rec = self.depth_decoder_n(rec)
61
62 output = (low_res_upsampled + rec) * (max_d - min_d) + min_d
63 return {dataset_keys.PREDICTION_DEPTH_IMG: output}

Callers 1

get_networkFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected