summaryrefslogtreecommitdiff
path: root/worldalign/vg_extract_vision.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/vg_extract_vision.py')
-rw-r--r--worldalign/vg_extract_vision.py150
1 files changed, 150 insertions, 0 deletions
diff --git a/worldalign/vg_extract_vision.py b/worldalign/vg_extract_vision.py
new file mode 100644
index 0000000..504840a
--- /dev/null
+++ b/worldalign/vg_extract_vision.py
@@ -0,0 +1,150 @@
+from __future__ import annotations
+
+import argparse
+from concurrent.futures import ThreadPoolExecutor, as_completed
+import json
+from pathlib import Path
+import time
+
+from PIL import Image
+import requests
+import torch
+from tqdm import tqdm
+from transformers import AutoImageProcessor, AutoModel
+
+from .common import dtype_for_device
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument(
+ "--nodes", default="artifacts/vg/vision_nodes.private.jsonl"
+ )
+ p.add_argument("--output", default="artifacts/vg/vision_features.pt")
+ p.add_argument(
+ "--image-cache", default="/tmp/yurenh2-worldalign-vg-images"
+ )
+ p.add_argument("--model", default="facebook/dinov2-small")
+ p.add_argument("--device", default="cuda:1")
+ p.add_argument("--node-batch-size", type=int, default=8)
+ p.add_argument("--download-workers", type=int, default=32)
+ p.add_argument("--limit", type=int)
+ return p.parse_args()
+
+
+def read_jsonl(path: str) -> list[dict]:
+ with open(path, encoding="utf-8") as handle:
+ return [json.loads(line) for line in handle if line.strip()]
+
+
+def download_one(record: dict, image_cache: Path) -> Path:
+ image_cache.mkdir(parents=True, exist_ok=True)
+ output = image_cache / f"{int(record['source_image_id'])}.jpg"
+ if output.exists() and output.stat().st_size > 0:
+ return output
+ temporary = output.with_suffix(".jpg.part")
+ last_error: Exception | None = None
+ for attempt in range(3):
+ try:
+ response = requests.get(record["url"], timeout=30)
+ response.raise_for_status()
+ temporary.write_bytes(response.content)
+ with Image.open(temporary) as image:
+ image.verify()
+ temporary.replace(output)
+ return output
+ except Exception as error:
+ last_error = error
+ if temporary.exists():
+ temporary.unlink()
+ time.sleep(1 + attempt)
+ raise RuntimeError(f"Failed to download {record['url']}") from last_error
+
+
+def crop_views(record: dict, path: Path) -> tuple[Image.Image, list[Image.Image]]:
+ with Image.open(path) as source:
+ image = source.convert("RGB")
+ width, height = image.size
+ crops: list[Image.Image] = []
+ for region in record["regions"]:
+ x0 = max(0, min(int(region["x"]), width - 1))
+ y0 = max(0, min(int(region["y"]), height - 1))
+ x1 = max(x0 + 1, min(x0 + int(region["width"]), width))
+ y1 = max(y0 + 1, min(y0 + int(region["height"]), height))
+ crops.append(image.crop((x0, y0, x1, y1)))
+ return image, crops
+
+
+@torch.inference_mode()
+def main() -> None:
+ args = parse_args()
+ records = read_jsonl(args.nodes)
+ if args.limit:
+ records = records[: args.limit]
+ if not records:
+ raise ValueError("No vision nodes found")
+ view_count = len(records[0]["regions"])
+ if any(len(record["regions"]) != view_count for record in records):
+ raise ValueError("All nodes must have the same number of region views")
+
+ image_cache = Path(args.image_cache)
+ paths: dict[str, Path] = {}
+ with ThreadPoolExecutor(max_workers=args.download_workers) as executor:
+ futures = {
+ executor.submit(download_one, record, image_cache): record["node_id"]
+ for record in records
+ }
+ for future in tqdm(
+ as_completed(futures), total=len(futures), desc="VG image download"
+ ):
+ paths[futures[future]] = future.result()
+
+ processor = AutoImageProcessor.from_pretrained(args.model)
+ dtype = dtype_for_device(args.device)
+ model = AutoModel.from_pretrained(args.model, torch_dtype=dtype).to(
+ args.device
+ )
+ model.eval()
+
+ region_outputs: list[torch.Tensor] = []
+ global_outputs: list[torch.Tensor] = []
+ for start in tqdm(
+ range(0, len(records), args.node_batch_size),
+ desc="DINO VG bundles",
+ ):
+ batch = records[start : start + args.node_batch_size]
+ images: list[Image.Image] = []
+ for record in batch:
+ global_image, crops = crop_views(record, paths[record["node_id"]])
+ images.append(global_image)
+ images.extend(crops)
+ pixels = processor(images=images, return_tensors="pt")[
+ "pixel_values"
+ ].to(args.device, dtype=dtype)
+ hidden = model(pixel_values=pixels, return_dict=True).last_hidden_state[
+ :, 0
+ ]
+ hidden = hidden.float().cpu().reshape(
+ len(batch), view_count + 1, -1
+ )
+ global_outputs.append(hidden[:, 0])
+ region_outputs.append(hidden[:, 1:])
+
+ state = {
+ "model": args.model,
+ "node_ids": [record["node_id"] for record in records],
+ "region_features": torch.cat(region_outputs),
+ "global_features": torch.cat(global_outputs),
+ "views_per_node": view_count,
+ "source_metadata_removed": True,
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(state, args.output)
+ print(
+ f"Wrote {args.output}: regions={tuple(state['region_features'].shape)}, "
+ f"global={tuple(state['global_features'].shape)}"
+ )
+
+
+if __name__ == "__main__":
+ main()