MCPcopy Create free account
hub / github.com/InternRobotics/G2VLM / align_points_scale_xyz_shift

Function align_points_scale_xyz_shift

modeling/pi3/utils/alignment.py:305–355  ·  view source on GitHub ↗

Align `points_src` to `points_tgt` with respect to a shared xyz scale and z shift. It is similar to `align_affine` but scale and shift are applied to different dimensions. ### Parameters: - `points_src: torch.Tensor` of shape (..., N, 3) - `points_tgt: torch.Tensor` of shape (

(points_src: torch.Tensor, points_tgt: torch.Tensor, weight: Optional[torch.Tensor], trunc: Optional[Union[float, torch.Tensor]] = None, max_iters: int = 30, eps: float = 1e-6)

Source from the content-addressed store, hash-verified

303
304
305def align_points_scale_xyz_shift(points_src: torch.Tensor, points_tgt: torch.Tensor, weight: Optional[torch.Tensor], trunc: Optional[Union[float, torch.Tensor]] = None, max_iters: int = 30, eps: float = 1e-6):
306 """
307 Align `points_src` to `points_tgt` with respect to a shared xyz scale and z shift.
308 It is similar to `align_affine` but scale and shift are applied to different dimensions.
309
310 ### Parameters:
311 - `points_src: torch.Tensor` of shape (..., N, 3)
312 - `points_tgt: torch.Tensor` of shape (..., N, 3)
313 - `weights: torch.Tensor` of shape (..., N)
314
315 ### Returns:
316 - `scale: torch.Tensor` of shape (...).
317 - `shift: torch.Tensor` of shape (..., 3)
318 """
319 dtype, device = points_src.dtype, points_src.device
320
321 # Flatten batch dimensions for simplicity
322 batch_shape, n = points_src.shape[:-2], points_src.shape[-2]
323 batch_size = math.prod(batch_shape)
324 points_src, points_tgt, weight = points_src.reshape(batch_size, n, 3), points_tgt.reshape(batch_size, n, 3), weight.reshape(batch_size, n)
325
326 # Take anchors
327 anchor_where_batch, anchor_where_n = torch.where(weight > 0)
328
329 with torch.no_grad():
330 points_src_anchor = points_src[anchor_where_batch, anchor_where_n] # (anchors, 3)
331 points_tgt_anchor = points_tgt[anchor_where_batch, anchor_where_n] # (anchors, 3)
332
333 points_src_anchored = points_src[anchor_where_batch, :, :] - points_src_anchor[..., None, :] # (anchors, n, 3)
334 points_tgt_anchored = points_tgt[anchor_where_batch, :, :] - points_tgt_anchor[..., None, :] # (anchors, n, 3)
335 weight_anchored = weight[anchor_where_batch, :, None].expand(-1, -1, 3) # (anchors, n, 3)
336
337 # Solve optimal scale and shift for each anchor
338 MAX_ELEMENTS = 2 ** 20
339 scale, loss, index = split_batch_fwd(align, MAX_ELEMENTS // 2, points_src_anchored.flatten(-2), points_tgt_anchored.flatten(-2), weight_anchored.flatten(-2), trunc) # (anchors,)
340
341 # Get optimal scale and shift for each batch element
342 loss, index_anchor = scatter_min(size=batch_size, dim=0, index=anchor_where_batch, src=loss) # (batch_size,)
343
344 index_2 = index[index_anchor] # (batch_size,) [0, 3n)
345 index_1 = anchor_where_n[index_anchor] * 3 + index_2 % 3 # (batch_size,) [0, 3n)
346
347 src_1, tgt_1 = torch.gather(points_src.flatten(-2), dim=1, index=index_1[..., None]).squeeze(-1), torch.gather(points_tgt.flatten(-2), dim=1, index=index_1[..., None]).squeeze(-1)
348 src_2, tgt_2 = torch.gather(points_src.flatten(-2), dim=1, index=index_2[..., None]).squeeze(-1), torch.gather(points_tgt.flatten(-2), dim=1, index=index_2[..., None]).squeeze(-1)
349
350 scale = (tgt_2 - tgt_1) / torch.where(src_2 != src_1, src_2 - src_1, 1.0)
351 shift = torch.gather(points_tgt, dim=1, index=(index_1 // 3)[..., None, None].expand(-1, -1, 3)).squeeze(-2) - scale[..., None] * torch.gather(points_src, dim=1, index=(index_1 // 3)[..., None, None].expand(-1, -1, 3)).squeeze(-2)
352
353 scale, shift = scale.reshape(batch_shape), shift.reshape(*batch_shape, 3)
354
355 return scale, shift
356
357
358def align_points_z_shift(points_src: torch.Tensor, points_tgt: torch.Tensor, weight: Optional[torch.Tensor], trunc: Optional[Union[float, torch.Tensor]] = None, max_iters: int = 30, eps: float = 1e-6):

Callers

nothing calls this directly

Calls 2

split_batch_fwdFunction · 0.70
scatter_minFunction · 0.70

Tested by

no test coverage detected