MP_PDE / scripts /result.py
yushuang88's picture
Upload folder using huggingface_hub
c92f17c verified
Raw
History Blame Contribute Delete
7.14 kB
"""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()