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