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