File size: 1,780 Bytes
3571a70
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Evaluate conditional token generation and visualize patch predictions."""

import json
from pathlib import Path

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import yaml


ROOT = Path(__file__).resolve().parents[1]


def main():
    config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
    data = np.load(ROOT / config["paths"]["inference_dir"] / "predictions.npz")
    metrics = {"samples": int(len(data["embedding"])), "mean_embedding_norm": float(np.linalg.norm(data["embedding"], axis=1).mean()),
               "conditioning_modalities": data["conditioning_modalities"].tolist(), "token_accuracy": {},
               "random_token_accuracy": 1.0 / int(config["model"]["engineering_vocab_size"])}
    for name in ("lulc", "ndvi", "s1grd"):
        metrics["token_accuracy"][name] = float((data[f"generated_{name}"] == data[f"target_{name}"]).mean())
    output = ROOT / config["paths"]["evaluation_dir"]
    output.mkdir(parents=True, exist_ok=True)
    (output / "metrics.json").write_text(json.dumps(metrics, indent=2) + "\n")
    figure, axes = plt.subplots(2, 3, figsize=(9, 6))
    grid = int(np.sqrt(data["generated_lulc"].shape[1]))
    for column, name in enumerate(("lulc", "ndvi", "s1grd")):
        axes[0, column].imshow(data[f"target_{name}"][0].reshape(grid, grid), cmap="viridis")
        axes[1, column].imshow(data[f"generated_{name}"][0].reshape(grid, grid), cmap="viridis")
        axes[0, column].set_title(f"{name} target")
        axes[1, column].set_title(f"{name} generated")
        axes[0, column].axis("off")
        axes[1, column].axis("off")
    figure.tight_layout()
    figure.savefig(output / "comparison.png", dpi=150)
    plt.close(figure)


if __name__ == "__main__":
    main()