File size: 2,149 Bytes
2ca760b
7002f4e
2ca760b
7002f4e
 
 
 
 
 
 
 
 
2ca760b
 
 
7002f4e
 
 
2ca760b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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()