Download scripts/result.py from OneScience-Group/RemoteCLIP: direct link, hf CLI and curl.
- Browser
- Download file 2.15 kB
-
https://huggingface.co/OneScience-Group/RemoteCLIP/resolve/main/scripts/result.py
- Command line
-
hf download hf://OneScience-Group/RemoteCLIP/scripts/result.py
-
curl -L -o result.py https://huggingface.co/OneScience-Group/RemoteCLIP/resolve/main/scripts/result.py
2.15 kB
| """Evaluate multi-positive bidirectional retrieval and visualize similarities.""" | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| import yaml | |
| ROOT = Path(__file__).resolve().parents[1] | |
| def recall(scores, query_ids, candidate_ids, k): | |
| top = np.argsort(-scores, axis=1)[:, :min(k, scores.shape[1])] | |
| return float(np.mean([np.isin(candidate_ids[index], query_ids[row]).any() for row, index in enumerate(top)])) | |
| def main(): | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml") | |
| parser.add_argument("--input", type=Path); parser.add_argument("--output-dir", type=Path) | |
| args = parser.parse_args(); config = yaml.safe_load(args.config.read_text()) | |
| source = args.input or ROOT / config["paths"]["inference_dir"] / "retrieval.npz" | |
| if not source.is_file(): raise FileNotFoundError("Run inference before evaluation") | |
| archive = np.load(source); scores, ids = archive["similarities"], archive["pair_ids"] | |
| metrics = {} | |
| for k in (1, 5, 10): | |
| metrics[f"image_to_text_R@{k}"] = recall(scores, ids, ids, k) | |
| metrics[f"text_to_image_R@{k}"] = recall(scores.T, ids, ids, k) | |
| metrics["mean_recall"] = float(np.mean(list(metrics.values()))) | |
| metrics.update(samples=int(len(ids)), protocol=str(archive["protocol"]), checkpoint=str(archive["checkpoint"]), | |
| multi_positive=True, data_source=str(archive["data_source"])) | |
| output = args.output_dir or ROOT / config["paths"]["evaluation_dir"]; output.mkdir(parents=True, exist_ok=True) | |
| (output / "metrics.json").write_text(json.dumps(metrics, indent=2) + "\n") | |
| figure, axis = plt.subplots(figsize=(5.4, 4.5)); image = axis.imshow(scores, cmap="magma") | |
| axis.set(xlabel="Text candidate", ylabel="Image query", title="RemoteCLIP cosine similarity") | |
| figure.colorbar(image, ax=axis); figure.tight_layout(); figure.savefig(output / "similarity_matrix.png", dpi=160); plt.close(figure) | |
| print(json.dumps(metrics, indent=2)); print(f"evaluation={output}") | |
| if __name__ == "__main__": main() | |