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