summaryrefslogtreecommitdiff
path: root/worldalign/extract_vision.py
blob: 348c6a49bcf11ae0a2bb8833df17cd06ca809ddb (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
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()