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()