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