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