Inspired by: Depth Map Super-Resolution by Deep Multi-Scale Guidance http://personal.ie.cuhk.edu.hk/~ccloy/files/eccv_2016_depth.pdf
| 8 | |
| 9 | |
| 10 | class 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} |