renorm pointmaps pts1, pts2 with norm_mode
(pts1, pts2, norm_mode='avg_dis', valid1=None, valid2=None, ret_factor=False)
| 247 | |
| 248 | |
| 249 | def normalize_pointcloud(pts1, pts2, norm_mode='avg_dis', valid1=None, valid2=None, ret_factor=False): |
| 250 | """ renorm pointmaps pts1, pts2 with norm_mode |
| 251 | """ |
| 252 | assert pts1.ndim >= 3 and pts1.shape[-1] == 3 |
| 253 | assert pts2 is None or (pts2.ndim >= 3 and pts2.shape[-1] == 3) |
| 254 | norm_mode, dis_mode = norm_mode.split('_') |
| 255 | |
| 256 | if norm_mode == 'avg': |
| 257 | # gather all points together (joint normalization) |
| 258 | nan_pts1, nnz1 = invalid_to_zeros(pts1, valid1, ndim=3) |
| 259 | nan_pts2, nnz2 = invalid_to_zeros(pts2, valid2, ndim=3) if pts2 is not None else (None, 0) |
| 260 | all_pts = torch.cat((nan_pts1, nan_pts2), dim=1) if pts2 is not None else nan_pts1 |
| 261 | |
| 262 | # compute distance to origin |
| 263 | all_dis = all_pts.norm(dim=-1) |
| 264 | if dis_mode == 'dis': |
| 265 | pass # do nothing |
| 266 | elif dis_mode == 'log1p': |
| 267 | all_dis = torch.log1p(all_dis) |
| 268 | elif dis_mode == 'warp-log1p': |
| 269 | # actually warp input points before normalizing them |
| 270 | log_dis = torch.log1p(all_dis) |
| 271 | warp_factor = log_dis / all_dis.clip(min=1e-8) |
| 272 | H1, W1 = pts1.shape[1:-1] |
| 273 | pts1 = pts1 * warp_factor[:, :W1 * H1].view(-1, H1, W1, 1) |
| 274 | if pts2 is not None: |
| 275 | H2, W2 = pts2.shape[1:-1] |
| 276 | pts2 = pts2 * warp_factor[:, W1 * H1:].view(-1, H2, W2, 1) |
| 277 | all_dis = log_dis # this is their true distance afterwards |
| 278 | else: |
| 279 | raise ValueError(f'bad {dis_mode=}') |
| 280 | |
| 281 | norm_factor = all_dis.sum(dim=1) / (nnz1 + nnz2 + 1e-8) |
| 282 | else: |
| 283 | # gather all points together (joint normalization) |
| 284 | nan_pts1 = invalid_to_nans(pts1, valid1, ndim=3) |
| 285 | nan_pts2 = invalid_to_nans(pts2, valid2, ndim=3) if pts2 is not None else None |
| 286 | all_pts = torch.cat((nan_pts1, nan_pts2), dim=1) if pts2 is not None else nan_pts1 |
| 287 | |
| 288 | # compute distance to origin |
| 289 | all_dis = all_pts.norm(dim=-1) |
| 290 | |
| 291 | if norm_mode == 'avg': |
| 292 | norm_factor = all_dis.nanmean(dim=1) |
| 293 | elif norm_mode == 'median': |
| 294 | norm_factor = all_dis.nanmedian(dim=1).values.detach() |
| 295 | elif norm_mode == 'sqrt': |
| 296 | norm_factor = all_dis.sqrt().nanmean(dim=1)**2 |
| 297 | else: |
| 298 | raise ValueError(f'bad {norm_mode=}') |
| 299 | |
| 300 | norm_factor = norm_factor.clip(min=1e-8) |
| 301 | while norm_factor.ndim < pts1.ndim: |
| 302 | norm_factor.unsqueeze_(-1) |
| 303 | |
| 304 | res = pts1 / norm_factor |
| 305 | if pts2 is not None: |
| 306 | res = (res, pts2 / norm_factor) |
no test coverage detected