File size: 3,079 Bytes
2ca760b
7002f4e
2ca760b
7002f4e
 
 
 
 
 
 
 
 
 
2ca760b
 
 
 
 
 
 
 
 
 
 
 
 
7002f4e
2ca760b
 
 
 
 
 
 
7002f4e
2ca760b
 
 
 
 
 
 
 
 
 
 
 
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
"""Run checkpoint-backed RemoteCLIP retrieval inference."""

import argparse
import importlib.util
from pathlib import Path
import numpy as np
import torch
import yaml

ROOT = Path(__file__).resolve().parents[1]


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml")
    parser.add_argument("--data", type=Path); parser.add_argument("--checkpoint", type=Path)
    parser.add_argument("--output-dir", type=Path); parser.add_argument("--device", choices=("auto", "cpu", "cuda"), default="auto")
    args = parser.parse_args(); config = yaml.safe_load(args.config.read_text())
    checkpoint_path = args.checkpoint or ROOT / config["paths"]["checkpoint"]
    if not checkpoint_path.is_file(): raise FileNotFoundError(f"checkpoint not found: {checkpoint_path}")
    spec = importlib.util.spec_from_file_location("remoteclip", ROOT / "model/remoteclip.py")
    module = importlib.util.module_from_spec(spec); spec.loader.exec_module(module)
    checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False)
    model = module.RemoteCLIP(vocabulary_size=config["data"]["vocabulary_size"],
                              context_length=config["data"]["context_length"],
                              eot_token_id=config["data"]["eot_token_id"], **config["model"])
    model.load_state_dict(checkpoint["model"])
    use_cuda = torch.cuda.is_available() and args.device != "cpu"
    if args.device == "cuda" and not use_cuda: raise RuntimeError("CUDA requested but unavailable")
    device = torch.device("cuda" if use_cuda else "cpu"); model.to(device).eval()
    data_path = args.data or ROOT / config["data"]["root"] / "test.npz"
    train_spec = importlib.util.spec_from_file_location("remoteclip_train", ROOT / "scripts/train.py")
    train_module = importlib.util.module_from_spec(train_spec); train_spec.loader.exec_module(train_module)
    dataset = train_module.PairDataset(data_path, config); archive = np.load(data_path)
    with torch.inference_mode():
        image_features = model.encode_image(torch.from_numpy(archive["images"]).to(device))
        text_features = model.encode_text(torch.from_numpy(archive["tokens"]).to(device))
    output_dir = args.output_dir or ROOT / config["paths"]["inference_dir"]; output_dir.mkdir(parents=True, exist_ok=True)
    np.savez_compressed(output_dir / "retrieval.npz", similarities=(image_features @ text_features.T).cpu().numpy(),
                        image_features=image_features.cpu().numpy(), text_features=text_features.cpu().numpy(),
                        pair_ids=archive["pair_ids"], images=archive["images"], checkpoint=np.asarray(str(checkpoint_path)),
                        data_source=archive["data_source"] if "data_source" in archive else np.asarray("provided"),
                        protocol=archive["protocol"] if "protocol" in archive else np.asarray("provided_npz"))
    print(f"inference={output_dir / 'retrieval.npz'} checkpoint={checkpoint_path}")


if __name__ == "__main__": main()