"""Evaluate Prithvi reconstruction and visualize temporal HLS samples.""" 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()) prediction = np.load(ROOT / config["paths"]["inference_dir"] / "predictions.npz") pixels, reconstruction = prediction["pixels"], prediction["reconstruction"] error = np.abs(reconstruction - pixels) per_frame = error.mean(axis=(0, 1, 3, 4)) embeddings = prediction["embedding"] metrics = { "samples": int(len(pixels)), "masked_patch_mse": float(prediction["masked_patch_mse"]), "reconstruction_mae": float(error.mean()), "per_frame_reconstruction_mae": [float(value) for value in per_frame], "mean_embedding_norm": float(np.linalg.norm(embeddings, axis=1).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(3, int(config["data"]["frames"]), figsize=(12, 8)) for frame in range(int(config["data"]["frames"])): source = pixels[0, [2, 1, 0], frame].transpose(1, 2, 0) rebuilt = reconstruction[0, [2, 1, 0], frame].transpose(1, 2, 0) low, high = np.percentile(source, (2, 98)) source = np.clip((source - low) / max(high - low, 1e-6), 0, 1) rebuilt = np.clip((rebuilt - low) / max(high - low, 1e-6), 0, 1) axes[0, frame].imshow(source) axes[1, frame].imshow(rebuilt) axes[2, frame].imshow(error[0, :, frame].mean(axis=0), cmap="magma") axes[0, frame].set_title(f"time {frame + 1}") for axis in axes[:, frame]: axis.axis("off") axes[0, 0].set_ylabel("input") axes[1, 0].set_ylabel("reconstruction") axes[2, 0].set_ylabel("absolute error") figure.tight_layout() figure.savefig(output / "comparison.png", dpi=150) plt.close(figure) print(f"metrics={output.relative_to(ROOT) / 'metrics.json'}") if __name__ == "__main__": main()