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)
| 303 | |
| 304 | |
| 305 | def 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 | |
| 358 | def 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): |
nothing calls this directly
no test coverage detected