Download scripts/result.py from OneScience-Group/Surya: direct link, hf CLI and curl.
- Browser
- Download file 4.19 kB
-
https://huggingface.co/OneScience-Group/Surya/resolve/main/scripts/result.py
- Command line
-
hf download hf://OneScience-Group/Surya/scripts/result.py
-
curl -L -o result.py https://huggingface.co/OneScience-Group/Surya/resolve/main/scripts/result.py
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() | |