summaryrefslogtreecommitdiff
path: root/worldalign/extract_vision.py
diff options
context:
space:
mode:
Diffstat (limited to 'worldalign/extract_vision.py')
-rw-r--r--worldalign/extract_vision.py61
1 files changed, 61 insertions, 0 deletions
diff --git a/worldalign/extract_vision.py b/worldalign/extract_vision.py
new file mode 100644
index 0000000..348c6a4
--- /dev/null
+++ b/worldalign/extract_vision.py
@@ -0,0 +1,61 @@
+from __future__ import annotations
+
+import argparse
+from pathlib import Path
+
+import torch
+from datasets import load_dataset
+from tqdm import tqdm
+from transformers import AutoImageProcessor, AutoModel
+
+from .common import batch_indices, dtype_for_device, read_json
+
+
+def parse_args() -> argparse.Namespace:
+ p = argparse.ArgumentParser()
+ p.add_argument("--manifest", default="artifacts/manifest.json")
+ p.add_argument("--output", default="artifacts/vision.pt")
+ p.add_argument("--model", default="facebook/dinov2-small")
+ p.add_argument("--device", default="cuda:1")
+ p.add_argument("--batch-size", type=int, default=96)
+ p.add_argument("--limit", type=int)
+ return p.parse_args()
+
+
+@torch.inference_mode()
+def main() -> None:
+ args = parse_args()
+ manifest = read_json(args.manifest)
+ rows = manifest["all_rows"]
+ if args.limit:
+ rows = rows[: args.limit]
+ dataset = load_dataset(manifest["dataset"], split=manifest["dataset_split"])
+ 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()
+
+ outputs: list[torch.Tensor] = []
+ for ids in tqdm(
+ batch_indices(len(rows), args.batch_size), desc="DINO image features"
+ ):
+ images = [dataset[int(rows[i])]["image"].convert("RGB") for i in ids]
+ batch = processor(images=images, return_tensors="pt")
+ batch = {k: v.to(args.device) for k, v in batch.items()}
+ result = model(**batch)
+ feature = result.last_hidden_state[:, 0]
+ outputs.append(feature.float().cpu())
+
+ value = {
+ "model": args.model,
+ "rows": rows,
+ "features": torch.cat(outputs),
+ }
+ Path(args.output).parent.mkdir(parents=True, exist_ok=True)
+ torch.save(value, args.output)
+ print(f"Wrote {args.output}: {tuple(value['features'].shape)}")
+
+
+if __name__ == "__main__":
+ main()
+