| 362 | self.y = y |
| 363 | |
| 364 | def _process(self, data): |
| 365 | # Compute the xy coordinates in the tiling grid, for each point |
| 366 | xy = data.pos[:, :2].clone().view(-1, 2) |
| 367 | xy -= xy.min(dim=0).values.view(1, 2) |
| 368 | xy /= xy.max(dim=0).values.view(1, 2) |
| 369 | xy = xy.clip(min=0, max=1) * self.tiling.view(1, 2) |
| 370 | xy = xy.long() |
| 371 | |
| 372 | # Select only the points in the desired tile |
| 373 | idx = torch.where((xy[:, 0] == self.x) & (xy[:, 1] == self.y))[0] |
| 374 | |
| 375 | return data.select(idx)[0] |
| 376 | |
| 377 | |
| 378 | class SampleRecursiveMainXYAxisTiling(Transform): |