File size: 7,141 Bytes
c92f17c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
"""Create MP-PDE E3 visualizations from real inference/training artifacts."""

from __future__ import annotations

import argparse
import json
from pathlib import Path
from typing import Any, Dict

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import yaml

PROJECT_ROOT = Path(__file__).resolve().parents[1]


def load_config(path: Path) -> Dict[str, Any]:
    with path.open("r", encoding="utf-8") as stream:
        return yaml.safe_load(stream)


def project_path(value: str | Path) -> Path:
    path = Path(value)
    return path if path.is_absolute() else PROJECT_ROOT / path


def main() -> None:
    parser = argparse.ArgumentParser(description="Plot real MP-PDE E3 rollout results")
    parser.add_argument("--config", type=Path, default=PROJECT_ROOT / "config/config.yaml")
    parser.add_argument("--predictions", type=Path)
    parser.add_argument("--metrics", type=Path)
    parser.add_argument("--history", type=Path)
    parser.add_argument("--output-dir", type=Path)
    parser.add_argument("--sample-index", type=int)
    args = parser.parse_args()
    config = load_config(args.config.resolve())
    predictions_path = args.predictions or project_path(config["paths"]["predictions"])
    metrics_path = args.metrics or project_path(config["paths"]["metrics"])
    history_path = args.history or project_path(config["paths"]["train_history"])
    output_dir = args.output_dir or project_path(config["paths"]["results"])
    if not predictions_path.is_file() or not metrics_path.is_file():
        raise FileNotFoundError(
            f"Real inference artifacts are required: {predictions_path} and {metrics_path}. "
            "Run scripts/inference.py after training; placeholder data will not be generated."
        )
    with np.load(predictions_path, allow_pickle=False) as archive:
        required = {"prediction", "target", "x", "t", "params", "sample_indices", "forecast_start_index", "per_time_mse"}
        missing = required.difference(archive.files)
        if missing:
            raise KeyError(f"predictions.npz is missing fields: {sorted(missing)}")
        prediction, target = archive["prediction"], archive["target"]
        x, t, per_time_mse = archive["x"], archive["t"], archive["per_time_mse"]
        forecast_start = int(archive["forecast_start_index"])
    with metrics_path.open("r", encoding="utf-8") as stream:
        metrics = json.load(stream)
    if prediction.shape != target.shape or prediction.ndim != 3:
        raise ValueError(f"Expected matching [S,T,N] arrays, found {prediction.shape}/{target.shape}")
    if prediction.shape[1:] != (t.size, x.size) or not np.all(np.isfinite(prediction)) or not np.all(np.isfinite(target)):
        raise ValueError("Prediction axes do not match x/t or contain non-finite values")
    recomputed = np.mean((prediction[:, forecast_start:] - target[:, forecast_start:]) ** 2, axis=(0, 2))
    if per_time_mse.shape != recomputed.shape or not np.allclose(per_time_mse, recomputed, rtol=2e-5, atol=1e-8):
        raise ValueError("Stored per_time_mse is inconsistent with prediction and target")
    if not np.isclose(float(metrics["accumulated_mse"]), float(np.sum(recomputed)), rtol=2e-5, atol=1e-8):
        raise ValueError("metrics.json accumulated_mse is inconsistent with predictions.npz")

    output_dir.mkdir(parents=True, exist_ok=True)
    sample = int(args.sample_index if args.sample_index is not None else config["visualization"]["sample_index"])
    if sample < 0 or sample >= prediction.shape[0]:
        raise IndexError(f"sample_index={sample} outside [0,{prediction.shape[0]})")
    dpi = int(config["visualization"]["dpi"])
    time_indices = [int(index) for index in config["visualization"]["time_indices"]]
    if any(index < 0 or index >= t.size for index in time_indices):
        raise IndexError(f"Configured time_indices exceed nt={t.size}")

    figure, axes = plt.subplots(len(time_indices), 1, figsize=(9, 2.4 * len(time_indices)), sharex=True)
    axes = np.atleast_1d(axes)
    for axis, time_index in zip(axes, time_indices):
        axis.plot(x, target[sample, time_index], color="black", linewidth=1.5, label="target")
        axis.plot(x, prediction[sample, time_index], color="tab:blue", linewidth=1.2, linestyle="--", label="MP-PDE")
        axis.set_ylabel("u")
        axis.set_title(f"t={t[time_index]:.4f}, index={time_index}")
        axis.grid(alpha=0.2)
    axes[0].legend(loc="best")
    axes[-1].set_xlabel("x")
    figure.tight_layout()
    rollout_path = output_dir / "e3_rollout.png"
    figure.savefig(rollout_path, dpi=dpi)
    plt.close(figure)

    absolute_error = np.abs(prediction[sample] - target[sample])
    figure, axes = plt.subplots(2, 1, figsize=(10, 7), gridspec_kw={"height_ratios": [2.2, 1.0]})
    image = axes[0].imshow(
        absolute_error.T, origin="lower", aspect="auto", extent=(float(t[0]), float(t[-1]), float(x[0]), float(x[-1])), cmap="magma"
    )
    axes[0].axvline(float(t[forecast_start]), color="white", linestyle="--", linewidth=1.0, label="forecast start")
    axes[0].set_ylabel("x")
    axes[0].set_title("Absolute rollout error")
    axes[0].legend(loc="upper right")
    figure.colorbar(image, ax=axes[0], label="|prediction-target|")
    axes[1].plot(t[forecast_start:], per_time_mse, color="tab:red")
    axes[1].set_xlabel("t")
    axes[1].set_ylabel("MSE")
    axes[1].set_title(f"Per-time MSE; accumulated={metrics['accumulated_mse']:.6g}")
    axes[1].grid(alpha=0.2)
    figure.tight_layout()
    error_path = output_dir / "e3_error.png"
    figure.savefig(error_path, dpi=dpi)
    plt.close(figure)

    created = [rollout_path, error_path]
    if history_path.is_file():
        with history_path.open("r", encoding="utf-8") as stream:
            history = json.load(stream)
        if not isinstance(history, list) or not history:
            raise ValueError(f"Training history is empty or malformed: {history_path}")
        epochs = [int(item["epoch"]) + 1 for item in history]
        figure, axes = plt.subplots(1, 2, figsize=(10, 4))
        axes[0].plot(epochs, [item["train_rmse"] for item in history], label="train bundle RMSE")
        axes[0].plot(epochs, [item["validation_bundle_rmse"] for item in history], label="validation bundle RMSE")
        axes[0].set_yscale("log")
        axes[0].set_xlabel("epoch")
        axes[0].set_ylabel("RMSE")
        axes[0].legend()
        axes[0].grid(alpha=0.2)
        axes[1].plot(epochs, [item["validation_accumulated_mse"] for item in history], color="tab:purple")
        axes[1].set_yscale("log")
        axes[1].set_xlabel("epoch")
        axes[1].set_ylabel("validation accumulated MSE")
        axes[1].grid(alpha=0.2)
        figure.tight_layout()
        training_path = output_dir / "training_curve.png"
        figure.savefig(training_path, dpi=dpi)
        plt.close(figure)
        created.append(training_path)
    else:
        print(f"Training history not found; skipped training curve: {history_path}", flush=True)
    for path in created:
        print(f"Saved figure: {path}", flush=True)


if __name__ == "__main__":
    main()