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()
|