diff options
Diffstat (limited to 'worldalign/natural_objects.py')
| -rw-r--r-- | worldalign/natural_objects.py | 47 |
1 files changed, 34 insertions, 13 deletions
diff --git a/worldalign/natural_objects.py b/worldalign/natural_objects.py index 65e035b..6e18259 100644 --- a/worldalign/natural_objects.py +++ b/worldalign/natural_objects.py @@ -62,7 +62,8 @@ def parse_args() -> argparse.Namespace: def spectral_segments( - features: torch.Tensor, grid: int, segments: int, seed: int + features: torch.Tensor, grid: int, segments: int, seed: int, + device: torch.device | str = "cpu", ) -> np.ndarray: """Segment a patch grid by clustering the normalised affinity spectrum. @@ -70,22 +71,42 @@ def spectral_segments( stay connected, then the leading eigenvectors of the normalised Laplacian are clustered. This is the standard unsupervised recipe over self-supervised features. + + The eigendecomposition dominates the run: one dense symmetric problem of + side grid^2 per image, which on a contended CPU costs more than the + forward pass that produced the features. Running it on the accelerator + that is already holding the model turns a multi-hour extraction into a + short one, and the spatial prior is built once and cached rather than + rebuilt per image. """ from sklearn.cluster import KMeans - normalised = F.normalize(features.double(), dim=-1) - affinity = (normalised @ normalised.T).clamp_min(0).numpy() - coordinates = np.stack( - np.meshgrid(np.arange(grid), np.arange(grid), indexing="ij"), -1 - ).reshape(-1, 2).astype(np.float64) - distance = ((coordinates[:, None, :] - coordinates[None, :, :]) ** 2).sum(-1) - affinity = affinity * np.exp(-distance / (2 * (grid / 4.0) ** 2)) + normalised = F.normalize(features.to(device).double(), dim=-1) + affinity = (normalised @ normalised.T).clamp_min(0) + affinity = affinity * _spatial_prior(grid, device) degree = affinity.sum(1) - laplacian = affinity / np.sqrt(np.outer(degree, degree) + 1e-9) - values, vectors = np.linalg.eigh(laplacian) + laplacian = affinity / torch.sqrt(torch.outer(degree, degree) + 1e-9) + values, vectors = torch.linalg.eigh(laplacian) embedding = vectors[:, -segments:] - embedding /= np.linalg.norm(embedding, axis=1, keepdims=True).clip(1e-9) - return KMeans(segments, n_init=10, random_state=seed).fit_predict(embedding) + embedding = embedding / embedding.norm(dim=1, keepdim=True).clamp_min(1e-9) + return KMeans(segments, n_init=10, random_state=seed).fit_predict( + embedding.cpu().numpy() + ) + + +_PRIOR_CACHE: dict[tuple, torch.Tensor] = {} + + +def _spatial_prior(grid: int, device: torch.device | str) -> torch.Tensor: + """Gaussian locality weight on the patch lattice, built once per grid.""" + key = (grid, str(device)) + if key not in _PRIOR_CACHE: + axis = torch.arange(grid, dtype=torch.float64, device=device) + rows, cols = torch.meshgrid(axis, axis, indexing="ij") + coordinates = torch.stack([rows.flatten(), cols.flatten()], dim=-1) + distance = torch.cdist(coordinates, coordinates) ** 2 + _PRIOR_CACHE[key] = torch.exp(-distance / (2 * (grid / 4.0) ** 2)) + return _PRIOR_CACHE[key] def describe_segments( @@ -220,7 +241,7 @@ def main() -> None: ) else: labels = spectral_segments( - patches[position], grid, args.segments, args.seed + patches[position], grid, args.segments, args.seed, args.device ) described = describe_segments( labels, patches[position], raw[position], grid, args.min_patches |
