(self, batch)
| 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} |
nothing calls this directly
no outgoing calls
no test coverage detected