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