RemoteCLIP / scripts /inference.py
zhangrenchao's picture
Add RemoteCLIP paper reproduction package
7002f4e verified
Raw
History Blame
2.44 kB
"""Extract RemoteCLIP features and image-text similarity scores."""
import importlib.util
from pathlib import Path
import numpy as np
import torch
import yaml
ROOT = Path(__file__).resolve().parents[1]
def load_model_class():
spec = importlib.util.spec_from_file_location("remoteclip_model", ROOT / "model" / "remoteclip.py")
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module.RemoteCLIP
def main():
with (ROOT / "conf" / "config.yaml").open(encoding="utf-8") as handle:
config = yaml.safe_load(handle)
checkpoint_path = ROOT / config["paths"]["checkpoint"]
data_path = ROOT / config["data"]["path"]
if not checkpoint_path.exists():
raise FileNotFoundError("Missing checkpoint. Run `python scripts/train.py` first.")
if not data_path.exists():
raise FileNotFoundError("Missing data. Run `python scripts/fake_data.py` first.")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)
RemoteCLIP = load_model_class()
model = RemoteCLIP(
vocabulary_size=config["data"]["vocabulary_size"],
context_length=config["data"]["context_length"],
**config["model"],
).to(device)
model.load_state_dict(checkpoint["model"])
model.eval()
archive = np.load(data_path)
images = torch.from_numpy(archive["test_images"]).to(device)
tokens = torch.from_numpy(archive["test_tokens"]).to(device)
with torch.inference_mode():
image_features = model.encode_image(images)
text_features = model.encode_text(tokens)
similarities = image_features @ text_features.t()
output_dir = ROOT / config["paths"]["inference_dir"]
output_dir.mkdir(parents=True, exist_ok=True)
np.savez_compressed(
output_dir / "retrieval.npz",
similarities=similarities.cpu().numpy(),
image_features=image_features.cpu().numpy(),
text_features=text_features.cpu().numpy(),
labels=archive["test_labels"],
images=archive["test_images"],
data_source=archive["data_source"],
protocol=archive["protocol"],
)
print(
f"output={output_dir.relative_to(ROOT)} samples={len(images)} "
f"data_source={str(archive['data_source'])} protocol={str(archive['protocol'])}"
)
if __name__ == "__main__":
main()