File size: 2,342 Bytes
7002f4e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 | """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()
|