| """Evaluate bidirectional retrieval and render a similarity matrix.""" |
|
|
| 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(similarities, k, transpose=False): |
| scores = similarities.T if transpose else similarities |
| topk = np.argsort(-scores, axis=1)[:, :k] |
| targets = np.arange(len(scores))[:, None] |
| return float((topk == targets).any(axis=1).mean()) |
|
|
|
|
| def main(): |
| with (ROOT / "conf" / "config.yaml").open(encoding="utf-8") as handle: |
| config = yaml.safe_load(handle) |
| input_path = ROOT / config["paths"]["inference_dir"] / "retrieval.npz" |
| if not input_path.exists(): |
| raise FileNotFoundError("Missing inference output. Run `python scripts/inference.py` first.") |
| archive = np.load(input_path) |
| similarities = archive["similarities"] |
| limit = len(similarities) |
| metrics = { |
| "image_to_text_r1": recall(similarities, 1), |
| "image_to_text_r5": recall(similarities, min(5, limit)), |
| "text_to_image_r1": recall(similarities, 1, True), |
| "text_to_image_r5": recall(similarities, min(5, limit), True), |
| "mean_recall": 0.0, |
| "samples": limit, |
| "data_source": str(archive["data_source"]), |
| "protocol": str(archive["protocol"]), |
| } |
| metrics["mean_recall"] = float( |
| np.mean( |
| [ |
| metrics["image_to_text_r1"], |
| metrics["image_to_text_r5"], |
| metrics["text_to_image_r1"], |
| metrics["text_to_image_r5"], |
| ] |
| ) |
| ) |
| output_dir = ROOT / config["paths"]["evaluation_dir"] |
| output_dir.mkdir(parents=True, exist_ok=True) |
| (output_dir / "metrics.json").write_text( |
| json.dumps(metrics, indent=2) + "\n", encoding="utf-8" |
| ) |
| figure, axis = plt.subplots(figsize=(5, 4)) |
| image = axis.imshow(similarities, cmap="viridis") |
| axis.set_xlabel("Text index") |
| axis.set_ylabel("Image index") |
| axis.set_title("RemoteCLIP image-text similarity") |
| figure.colorbar(image, ax=axis) |
| figure.tight_layout() |
| figure.savefig(output_dir / "similarity_matrix.png", dpi=120) |
| plt.close(figure) |
| print(json.dumps(metrics)) |
| print(f"evaluation={output_dir.relative_to(ROOT)}") |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|