Apply a geometric transformation to a list of 3-D points. H: 3x3 or 4x4 projection matrix (typically a Homography) p: numpy/torch/tuple of coordinates. Shape must be (...,2) or (...,3) ncol: int. number of columns of the result (2 or 3) norm: float. if != 0, the resut is projected
(Trf, pts, ncol=None, norm=False)
| 262 | |
| 263 | |
| 264 | def geotrf(Trf, pts, ncol=None, norm=False): |
| 265 | """Apply a geometric transformation to a list of 3-D points. |
| 266 | |
| 267 | H: 3x3 or 4x4 projection matrix (typically a Homography) |
| 268 | p: numpy/torch/tuple of coordinates. Shape must be (...,2) or (...,3) |
| 269 | |
| 270 | ncol: int. number of columns of the result (2 or 3) |
| 271 | norm: float. if != 0, the resut is projected on the z=norm plane. |
| 272 | |
| 273 | Returns an array of projected 2d points. |
| 274 | """ |
| 275 | assert Trf.ndim >= 2 |
| 276 | if isinstance(Trf, np.ndarray): |
| 277 | pts = np.asarray(pts) |
| 278 | elif isinstance(Trf, torch.Tensor): |
| 279 | pts = torch.as_tensor(pts, dtype=Trf.dtype) |
| 280 | |
| 281 | # adapt shape if necessary |
| 282 | output_reshape = pts.shape[:-1] |
| 283 | ncol = ncol or pts.shape[-1] |
| 284 | |
| 285 | # optimized code |
| 286 | if ( |
| 287 | isinstance(Trf, torch.Tensor) |
| 288 | and isinstance(pts, torch.Tensor) |
| 289 | and Trf.ndim == 3 |
| 290 | and pts.ndim == 4 |
| 291 | ): |
| 292 | d = pts.shape[3] |
| 293 | if Trf.shape[-1] == d: |
| 294 | pts = torch.einsum("bij, bhwj -> bhwi", Trf, pts) |
| 295 | elif Trf.shape[-1] == d + 1: |
| 296 | pts = ( |
| 297 | torch.einsum("bij, bhwj -> bhwi", Trf[:, :d, :d], pts) |
| 298 | + Trf[:, None, None, :d, d] |
| 299 | ) |
| 300 | else: |
| 301 | raise ValueError(f"bad shape, not ending with 3 or 4, for {pts.shape=}") |
| 302 | else: |
| 303 | if Trf.ndim >= 3: |
| 304 | n = Trf.ndim - 2 |
| 305 | assert Trf.shape[:n] == pts.shape[:n], "batch size does not match" |
| 306 | Trf = Trf.reshape(-1, Trf.shape[-2], Trf.shape[-1]) |
| 307 | |
| 308 | if pts.ndim > Trf.ndim: |
| 309 | # Trf == (B,d,d) & pts == (B,H,W,d) --> (B, H*W, d) |
| 310 | pts = pts.reshape(Trf.shape[0], -1, pts.shape[-1]) |
| 311 | elif pts.ndim == 2: |
| 312 | # Trf == (B,d,d) & pts == (B,d) --> (B, 1, d) |
| 313 | pts = pts[:, None, :] |
| 314 | |
| 315 | if pts.shape[-1] + 1 == Trf.shape[-1]: |
| 316 | Trf = Trf.swapaxes(-1, -2) # transpose Trf |
| 317 | pts = pts @ Trf[..., :-1, :] + Trf[..., -1:, :] |
| 318 | elif pts.shape[-1] == Trf.shape[-1]: |
| 319 | Trf = Trf.swapaxes(-1, -2) # transpose Trf |
| 320 | pts = pts @ Trf |
| 321 | else: |
no outgoing calls
no test coverage detected