Surya / scripts /result.py
zhangrenchao's picture
Upload Surya model package
a13f4b9 verified
Raw
History Blame Contribute Delete
4.19 kB
"""Evaluate Surya rollouts and generate solar forecasting figures."""
import json
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
import yaml
ROOT = Path(__file__).resolve().parents[1]
def main():
cfg = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
source = ROOT / cfg["paths"]["inference_dir"] / "forecast.npz"
if not source.exists(): raise FileNotFoundError("Run inference before evaluation")
data = np.load(source); targets, predictions = data["targets"], data["predictions"]
step_mse = np.mean((predictions - targets) ** 2, axis=(0, 2, 3, 4))
persistence = np.repeat(data["inputs"][:, -1:, :, :, :], targets.shape[1], axis=1)
persistence_mse = np.mean((persistence - targets) ** 2, axis=(0, 2, 3, 4))
skill = 1.0 - step_mse / np.maximum(persistence_mse, 1e-8)
channel_mse = np.mean((predictions - targets) ** 2, axis=(0, 1, 3, 4))
predicted_activity = predictions[:, :, :8].sum(axis=(2, 3, 4))
target_activity = targets[:, :, :8].sum(axis=(2, 3, 4))
result = {"forecast_mse": float(step_mse.mean()), "step_mse": step_mse.tolist(),
"persistence_step_mse": persistence_mse.tolist(), "persistence_skill": skill.tolist(),
"channel_mse": channel_mse.tolist(), "aia_mse": float(channel_mse[:8].mean()),
"hmi_mse": float(channel_mse[8:].mean()),
"activity_mae": float(np.mean(np.abs(predicted_activity - target_activity))),
"data_source": "synthetic", "protocol": cfg["data"]["protocol"], "baseline": "persistence"}
output = ROOT / cfg["paths"]["evaluation_dir"]; output.mkdir(parents=True, exist_ok=True)
(output / "metrics.json").write_text(json.dumps(result, indent=2) + "\n")
steps = np.arange(1, len(step_mse) + 1)
figure, axis = plt.subplots(figsize=(6, 3.5)); axis.plot(steps, step_mse, marker="o", label="Surya")
axis.plot(steps, persistence_mse, marker="s", label="Persistence"); axis.set(xlabel="Forecast step (hour)", ylabel="MSE", title="Autoregressive Forecast Skill")
axis.legend(); axis.grid(alpha=0.25); figure.tight_layout(); figure.savefig(output / "rollout_forecast_skill.png", dpi=160); plt.close(figure)
figure, axes = plt.subplots(2, len(steps), figsize=(2.5 * len(steps), 5))
for index in range(len(steps)):
axes[0, index].imshow(targets[0, index, 0], cmap="inferno"); axes[0, index].set_title(f"Target +{index + 1}h")
axes[1, index].imshow(predictions[0, index, 0], cmap="inferno"); axes[1, index].set_title(f"Surya +{index + 1}h")
axes[0, index].axis("off"); axes[1, index].axis("off")
figure.tight_layout(); figure.savefig(output / "solar_dynamics_forecast.png", dpi=160); plt.close(figure)
figure, axis = plt.subplots(figsize=(6, 3.5)); axis.plot(steps, target_activity[0], marker="o", label="Ground truth")
axis.plot(steps, predicted_activity[0], marker="s", label="Surya"); axis.set(xlabel="Forecast step (hour)", ylabel="Integrated AIA activity", title="Solar Activity Evolution")
axis.legend(); axis.grid(alpha=0.25); figure.tight_layout(); figure.savefig(output / "solar_activity_evolution.png", dpi=160); plt.close(figure)
figure, axis = plt.subplots(figsize=(7, 3.5)); axis.bar(np.arange(len(channel_mse)), channel_mse, color="#287271")
channel_names = [f"AIA {i + 1}" for i in range(8)] + [f"HMI {i + 1}" for i in range(5)]
axis.set_xticks(np.arange(len(channel_mse)), channel_names, rotation=45, ha="right")
axis.set(xlabel="SDO channel", ylabel="MSE", title="AIA/HMI Channel Forecast Error")
figure.tight_layout(); figure.savefig(output / "sdo_channel_error.png", dpi=160); plt.close(figure)
figure, axis = plt.subplots(figsize=(6, 3.5)); axis.plot(steps, skill, marker="o", color="#7a5195")
axis.axhline(0.0, color="black", linewidth=0.8)
axis.set(xlabel="Forecast step (hour)", ylabel="Skill vs persistence", title="Persistence Skill by Lead Time")
axis.grid(alpha=0.25); figure.tight_layout(); figure.savefig(output / "persistence_skill.png", dpi=160); plt.close(figure)
print(json.dumps(result, indent=2)); print("evaluation=", output)
if __name__ == "__main__": main()