PrithviEO / scripts /result.py
zhangrenchao's picture
Add engineering reproduction package
4c4d99c verified
Raw
History Blame Contribute Delete
2.23 kB
"""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()