Optimize a global scene, given a list of pairwise observations. Graph node: images Graph edges: observations = (pred1, pred2)
| 14 | |
| 15 | |
| 16 | class PointCloudOptimizer(BasePCOptimizer): |
| 17 | """ Optimize a global scene, given a list of pairwise observations. |
| 18 | Graph node: images |
| 19 | Graph edges: observations = (pred1, pred2) |
| 20 | """ |
| 21 | |
| 22 | def __init__(self, *args, optimize_pp=False, focal_break=20, **kwargs): |
| 23 | super().__init__(*args, **kwargs) |
| 24 | |
| 25 | self.has_im_poses = True # by definition of this class |
| 26 | self.focal_break = focal_break |
| 27 | |
| 28 | # adding thing to optimize |
| 29 | self.im_depthmaps = nn.ParameterList(torch.randn(H, W)/10-3 for H, W in self.imshapes) # log(depth) |
| 30 | self.im_poses = nn.ParameterList(self.rand_pose(self.POSE_DIM) for _ in range(self.n_imgs)) # camera poses |
| 31 | self.im_focals = nn.ParameterList(torch.FloatTensor( |
| 32 | [self.focal_break*np.log(max(H, W))]) for H, W in self.imshapes) # camera intrinsics |
| 33 | self.im_pp = nn.ParameterList(torch.zeros((2,)) for _ in range(self.n_imgs)) # camera intrinsics |
| 34 | self.im_pp.requires_grad_(optimize_pp) |
| 35 | |
| 36 | self.imshape = self.imshapes[0] |
| 37 | im_areas = [h*w for h, w in self.imshapes] |
| 38 | self.max_area = max(im_areas) |
| 39 | |
| 40 | # adding thing to optimize |
| 41 | self.im_depthmaps = ParameterStack(self.im_depthmaps, is_param=True, fill=self.max_area) |
| 42 | self.im_poses = ParameterStack(self.im_poses, is_param=True) |
| 43 | self.im_focals = ParameterStack(self.im_focals, is_param=True) |
| 44 | self.im_pp = ParameterStack(self.im_pp, is_param=True) |
| 45 | self.register_buffer('_pp', torch.tensor([(w/2, h/2) for h, w in self.imshapes])) |
| 46 | self.register_buffer('_grid', ParameterStack( |
| 47 | [xy_grid(W, H, device=self.device) for H, W in self.imshapes], fill=self.max_area)) |
| 48 | |
| 49 | # pre-compute pixel weights |
| 50 | self.register_buffer('_weight_i', ParameterStack( |
| 51 | [self.conf_trf(self.conf_i[i_j]) for i_j in self.str_edges], fill=self.max_area)) |
| 52 | self.register_buffer('_weight_j', ParameterStack( |
| 53 | [self.conf_trf(self.conf_j[i_j]) for i_j in self.str_edges], fill=self.max_area)) |
| 54 | |
| 55 | # precompute aa |
| 56 | self.register_buffer('_stacked_pred_i', ParameterStack(self.pred_i, self.str_edges, fill=self.max_area)) |
| 57 | self.register_buffer('_stacked_pred_j', ParameterStack(self.pred_j, self.str_edges, fill=self.max_area)) |
| 58 | self.register_buffer('_ei', torch.tensor([i for i, j in self.edges])) |
| 59 | self.register_buffer('_ej', torch.tensor([j for i, j in self.edges])) |
| 60 | self.total_area_i = sum([im_areas[i] for i, j in self.edges]) |
| 61 | self.total_area_j = sum([im_areas[j] for i, j in self.edges]) |
| 62 | |
| 63 | def _check_all_imgs_are_selected(self, msk): |
| 64 | assert np.all(self._get_msk_indices(msk) == np.arange(self.n_imgs)), 'incomplete mask!' |
| 65 | |
| 66 | def preset_pose(self, known_poses, pose_msk=None): # cam-to-world |
| 67 | self._check_all_imgs_are_selected(pose_msk) |
| 68 | |
| 69 | if isinstance(known_poses, torch.Tensor) and known_poses.ndim == 2: |
| 70 | known_poses = [known_poses] |
| 71 | for idx, pose in zip(self._get_msk_indices(pose_msk), known_poses): |
| 72 | if self.verbose: |
| 73 | print(f' (setting pose #{idx} = {pose[:3,3]})') |