| 118 | return trans_points |
| 119 | |
| 120 | def transform_points(points, tform, points_scale=None): |
| 121 | points_2d = points[:,:,:2] |
| 122 | # -1 to 1 |
| 123 | |
| 124 | #'input points must use original range' |
| 125 | if points_scale: |
| 126 | assert points_scale[0]==points_scale[1] |
| 127 | points_2d = (points_2d*0.5 + 0.5)*points_scale[0] |
| 128 | # import ipdb; ipdb.set_trace() |
| 129 | # 0 - points_scale |
| 130 | |
| 131 | batch_size, n_points, _ = points.shape |
| 132 | trans_points_2d = torch.bmm( |
| 133 | torch.cat([points_2d, |
| 134 | torch.ones([batch_size, n_points, 1], |
| 135 | device=points.device, |
| 136 | dtype=points.dtype)], |
| 137 | dim=-1), |
| 138 | tform |
| 139 | ) |
| 140 | trans_points = torch.cat([trans_points_2d[:,:,:2], points[:,:,2:]], dim=-1) |
| 141 | return trans_points |