diff options
Diffstat (limited to 'worldalign/extract_vision.py')
| -rw-r--r-- | worldalign/extract_vision.py | 61 |
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() + |
