summaryrefslogtreecommitdiff
path: root/worldalign/natural_objects.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/natural_objects.py')
-rw-r--r--worldalign/natural_objects.py262
1 files changed, 262 insertions, 0 deletions
diff --git a/worldalign/natural_objects.py b/worldalign/natural_objects.py
new file mode 100644
index 0000000..65e035b
--- /dev/null
+++ b/worldalign/natural_objects.py
@@ -0,0 +1,262 @@
+"""Unsupervised object states for natural images.
+
+The synthetic world's objects came from watershed on a flat background,
+which real photographs do not offer. Self-supervised vision features do:
+DINOv2 patch descriptors segment objects without labels, which is what
+the deep-spectral line of work exploits. Each image is cut into segments
+by spectral clustering of the patch affinity graph, and each segment is
+described by observables that a text corpus also states -- mean colour,
+relative area, position, and its own feature centroid for later category
+clustering.
+
+Nothing here reads text, pairs, or region annotations. Boxes from the
+preprocessing file are not used; segmentation is derived from pixels.
+"""
+
+from __future__ import annotations
+
+import argparse
+import json
+from pathlib import Path
+
+import numpy as np
+import torch
+import torch.nn.functional as F
+from PIL import Image
+from tqdm import tqdm
+from transformers import AutoImageProcessor, AutoModel
+
+from .common import batch_indices, read_json, seed_everything, write_json
+
+
+def parse_args() -> argparse.Namespace:
+ parser = argparse.ArgumentParser()
+ parser.add_argument("--vg-dir", default="artifacts/vg_5k")
+ parser.add_argument("--image-cache", default="/tmp/yurenh2-worldalign-vg-images")
+ parser.add_argument("--model", default="facebook/dinov2-base")
+ parser.add_argument("--image-size", type=int, default=448)
+ parser.add_argument("--segments", type=int, default=6)
+ parser.add_argument("--min-patches", type=int, default=8)
+ parser.add_argument(
+ "--oracle-regions",
+ action="store_true",
+ help="Diagnostic upper bound: take segments from the annotated "
+ "region boxes instead of unsupervised segmentation, isolating "
+ "how much of the field deficit segmentation accounts for.",
+ )
+ parser.add_argument(
+ "--views",
+ type=int,
+ default=1,
+ help="Augmentation orbit size. Each view is a random resized crop, "
+ "so content is fixed and framing varies -- the natural-image "
+ "analogue of the synthetic world's re-renders.",
+ )
+ parser.add_argument("--crop-low", type=float, default=0.75)
+ parser.add_argument("--nodes", type=int, default=0)
+ parser.add_argument("--batch-size", type=int, default=16)
+ parser.add_argument("--device", default="cuda:3")
+ parser.add_argument("--seed", type=int, default=0)
+ parser.add_argument("--output", default="artifacts/vg_5k/natural_objects.pt")
+ return parser.parse_args()
+
+
+def spectral_segments(
+ features: torch.Tensor, grid: int, segments: int, seed: int
+) -> np.ndarray:
+ """Segment a patch grid by clustering the normalised affinity spectrum.
+
+ Feature cosine affinity is combined with a spatial prior so segments
+ stay connected, then the leading eigenvectors of the normalised
+ Laplacian are clustered. This is the standard unsupervised recipe over
+ self-supervised features.
+ """
+ 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))
+ degree = affinity.sum(1)
+ laplacian = affinity / np.sqrt(np.outer(degree, degree) + 1e-9)
+ values, vectors = np.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)
+
+
+def describe_segments(
+ labels: np.ndarray,
+ features: torch.Tensor,
+ pixels: torch.Tensor,
+ grid: int,
+ min_patches: int,
+) -> list[dict]:
+ """Observables of each segment: colour, area, position, feature centre."""
+ image = pixels.permute(1, 2, 0).numpy()
+ size = image.shape[0]
+ scale = size // grid
+ described = []
+ for label in np.unique(labels):
+ mask = labels == label
+ if mask.sum() < min_patches:
+ continue
+ rows, cols = np.nonzero(mask.reshape(grid, grid))
+ pixel_mask = np.zeros((size, size), dtype=bool)
+ for row, col in zip(rows, cols):
+ pixel_mask[
+ row * scale : (row + 1) * scale, col * scale : (col + 1) * scale
+ ] = True
+ described.append(
+ {
+ "rgb": image[pixel_mask].mean(0),
+ "area": float(mask.mean()),
+ "centre": np.array([rows.mean() / grid, cols.mean() / grid]),
+ "feature": features[mask].mean(0).numpy(),
+ }
+ )
+ return described
+
+
+def region_segments(
+ record: dict, features: torch.Tensor, pixels: torch.Tensor, grid: int
+) -> list[dict]:
+ """Segments taken from annotated boxes: an oracle for segmentation only."""
+ image = pixels.permute(1, 2, 0).numpy()
+ size = image.shape[0]
+ width, height = record["width"], record["height"]
+ described = []
+ for region in record["regions"]:
+ x0 = max(0.0, min(region["x"] / width, 1.0)) * grid
+ x1 = max(0.0, min((region["x"] + region["width"]) / width, 1.0)) * grid
+ y0 = max(0.0, min(region["y"] / height, 1.0)) * grid
+ y1 = max(0.0, min((region["y"] + region["height"]) / height, 1.0)) * grid
+ columns = np.arange(grid) + 0.5
+ inside_x = (columns >= x0) & (columns <= x1)
+ inside_y = (columns >= y0) & (columns <= y1)
+ mask = (inside_y[:, None] & inside_x[None, :]).flatten()
+ if not mask.any():
+ centre_x = min(grid - 1, max(0, int((x0 + x1) / 2)))
+ centre_y = min(grid - 1, max(0, int((y0 + y1) / 2)))
+ mask = np.zeros(grid * grid, dtype=bool)
+ mask[centre_y * grid + centre_x] = True
+ rows, cols = np.nonzero(mask.reshape(grid, grid))
+ scale = size // grid
+ pixel_mask = np.zeros((size, size), dtype=bool)
+ for row, col in zip(rows, cols):
+ pixel_mask[
+ row * scale : (row + 1) * scale, col * scale : (col + 1) * scale
+ ] = True
+ described.append(
+ {
+ "rgb": image[pixel_mask].mean(0),
+ "area": float(mask.mean()),
+ "centre": np.array([rows.mean() / grid, cols.mean() / grid]),
+ "feature": features[mask].mean(0).numpy(),
+ }
+ )
+ return described
+
+
+@torch.inference_mode()
+def main() -> None:
+ args = parse_args()
+ seed_everything(args.seed)
+ records = [
+ json.loads(line)
+ for line in Path(args.vg_dir, "vision_nodes.private.jsonl")
+ .read_text(encoding="utf-8")
+ .splitlines()
+ if line.strip()
+ ]
+ if args.nodes:
+ records = records[: args.nodes]
+ processor = AutoImageProcessor.from_pretrained(args.model)
+ model = AutoModel.from_pretrained(args.model, torch_dtype=torch.float32).to(
+ args.device
+ )
+ model.eval()
+ patch = model.config.patch_size
+ grid = args.image_size // patch
+ mean = torch.tensor(processor.image_mean).view(3, 1, 1)
+ std = torch.tensor(processor.image_std).view(3, 1, 1)
+
+ generator = np.random.default_rng(args.seed)
+
+ def load(record: dict, view: int = 0) -> torch.Tensor:
+ path = Path(args.image_cache, f"{record['source_image_id']}.jpg")
+ with Image.open(path) as image:
+ picture = image.convert("RGB")
+ if view > 0:
+ width, height = picture.size
+ scale = generator.uniform(args.crop_low, 1.0)
+ box_w, box_h = int(width * scale), int(height * scale)
+ left = int(generator.integers(0, max(width - box_w, 1)))
+ top = int(generator.integers(0, max(height - box_h, 1)))
+ picture = picture.crop((left, top, left + box_w, top + box_h))
+ resized = picture.resize(
+ (args.image_size, args.image_size), Image.BILINEAR
+ )
+ return torch.from_numpy(np.asarray(resized).copy()).permute(2, 0, 1).float() / 255.0
+
+ node_ids, all_segments = [], []
+ view_segments: list[list] = [[] for _ in range(args.views)]
+ for indices in tqdm(
+ list(batch_indices(len(records), args.batch_size)), desc="segment"
+ ):
+ for view in range(args.views):
+ batch = [records[index] for index in indices]
+ raw = torch.stack([load(record, view) for record in batch])
+ normalised = ((raw - mean) / std).to(args.device)
+ hidden = model(pixel_values=normalised, return_dict=True).last_hidden_state
+ patches = hidden[:, 1:].float().cpu()
+ for position, record in enumerate(batch):
+ if args.oracle_regions:
+ described = region_segments(
+ record, patches[position], raw[position], grid
+ )
+ else:
+ labels = spectral_segments(
+ patches[position], grid, args.segments, args.seed
+ )
+ described = describe_segments(
+ labels, patches[position], raw[position], grid, args.min_patches
+ )
+ if view == 0:
+ node_ids.append(record["node_id"])
+ all_segments.append(described)
+ view_segments[view].append(described)
+
+ torch.save(
+ {
+ "model": args.model,
+ "node_ids": node_ids,
+ "segments": all_segments,
+ "view_segments": view_segments if args.views > 1 else None,
+ "views": args.views,
+ "grid": grid,
+ "protocol": (
+ "Segments come from spectral clustering of self-supervised "
+ "patch features; no boxes, no text, no pairs."
+ ),
+ },
+ args.output,
+ )
+ counts = [len(item) for item in all_segments]
+ print(
+ json.dumps(
+ {
+ "nodes": len(node_ids),
+ "mean_segments": float(np.mean(counts)),
+ "min_segments": int(np.min(counts)),
+ }
+ )
+ )
+ print(f"Wrote {args.output}")
+
+
+if __name__ == "__main__":
+ main()