MCPcopy Create free account
hub / github.com/NVlabs/InstantSplat / normalize_pointcloud

Function normalize_pointcloud

dust3r/utils/geometry.py:249–309  ·  view source on GitHub ↗

renorm pointmaps pts1, pts2 with norm_mode

(pts1, pts2, norm_mode='avg_dis', valid1=None, valid2=None, ret_factor=False)

Source from the content-addressed store, hash-verified

247
248
249def 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)

Callers 2

get_all_pts3dMethod · 0.90
get_all_pts3dMethod · 0.90

Calls 2

invalid_to_zerosFunction · 0.90
invalid_to_nansFunction · 0.90

Tested by

no test coverage detected