RemoteCLIP / scripts /inference.py
zhangrenchao's picture
Update RemoteCLIP model package
2ca760b verified
Raw
History Blame Contribute Delete
3.08 kB
"""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()