File size: 4,188 Bytes
a13f4b9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 | """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()
|