Saves a PLY visualization of the world points and images. Args:
(
pred_dict, save_path, init_conf_threshold=20.0, filter_nan=True, verbose=True)
| 123 | |
| 124 | |
| 125 | def save_ply_visualization( |
| 126 | pred_dict, save_path, init_conf_threshold=20.0, filter_nan=True, verbose=True): |
| 127 | """ |
| 128 | Saves a PLY visualization of the world points and images. |
| 129 | Args: |
| 130 | """ |
| 131 | pred_dict = pred_dict.copy() |
| 132 | for key in pred_dict.keys(): |
| 133 | if isinstance(pred_dict[key], torch.Tensor): |
| 134 | pred_dict[key] = pred_dict[key].detach() |
| 135 | |
| 136 | index = 0 |
| 137 | seq_len = len(pred_dict["images"][index]) |
| 138 | |
| 139 | images = pred_dict["images"][index] |
| 140 | images_grid = torchvision.utils.make_grid( |
| 141 | images, nrow=8, normalize=False, scale_each=False |
| 142 | ) |
| 143 | images_grid = images_grid |
| 144 | images_save_path = f'input_images.png' |
| 145 | torchvision.utils.save_image( |
| 146 | images_grid, images_save_path, normalize=False, scale_each=False |
| 147 | ) |
| 148 | |
| 149 | pred_pts = pred_dict["points"][index] |
| 150 | |
| 151 | data_h, data_w = images.shape[-2:] |
| 152 | global_points = pred_pts |
| 153 | |
| 154 | data_size = (data_h, data_w) |
| 155 | global_points = F.interpolate( |
| 156 | global_points.permute(0, 3, 1, 2), data_size, |
| 157 | mode="bilinear", align_corners=False, antialias=True |
| 158 | ).permute(0, 2, 3, 1) |
| 159 | |
| 160 | global_points = global_points.cpu().numpy() |
| 161 | pred_pts = global_points |
| 162 | |
| 163 | colors = images.permute(0, 2, 3, 1).cpu().numpy().reshape(-1, 3) |
| 164 | |
| 165 | pred_pts = pred_pts.reshape(-1, 3) |
| 166 | |
| 167 | |
| 168 | if filter_nan: |
| 169 | nan_mask = np.any(np.isnan(pred_pts), axis=1) |
| 170 | inf_mask = np.any(np.isinf(pred_pts), axis=1) |
| 171 | invalid_mask = np.logical_or(nan_mask, inf_mask) |
| 172 | |
| 173 | valid_mask = ~invalid_mask |
| 174 | |
| 175 | total_points = len(pred_pts) |
| 176 | invalid_count = np.sum(invalid_mask) |
| 177 | valid_count = np.sum(valid_mask) |
| 178 | |
| 179 | if invalid_count > 0: |
| 180 | pred_pts_filtered = pred_pts[valid_mask] |
| 181 | colors_filtered = colors[valid_mask] |
| 182 |