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