SEEDS / scripts /result.py
yzt15806542928's picture
Upload folder using huggingface_hub
1ca0208 verified
Raw
History Blame Contribute Delete
3.97 kB
"""Evaluate and visualize generated SEEDS ensembles."""
from __future__ import annotations
import argparse
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
from common import load_config, resolve_path
def _crps(prediction: np.ndarray, target: np.ndarray) -> float:
first = np.mean(np.abs(prediction - target[:, None]), axis=1)
pairwise = np.mean(np.abs(prediction[:, :, None] - prediction[:, None, :]), axis=(1, 2))
return float(np.mean(first - 0.5 * pairwise))
def _acc(prediction: np.ndarray, target: np.ndarray) -> float:
forecast = prediction.mean(axis=1).reshape(prediction.shape[0], -1)
truth = target.reshape(target.shape[0], -1)
forecast = forecast - forecast.mean(axis=1, keepdims=True)
truth = truth - truth.mean(axis=1, keepdims=True)
numerator = np.sum(forecast * truth, axis=1)
denominator = np.sqrt(np.sum(forecast**2, axis=1) * np.sum(truth**2, axis=1))
return float(np.mean(numerator / np.maximum(denominator, 1e-8)))
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--config", default="conf/config.yaml")
parser.add_argument("--prediction", default=None)
parser.add_argument("--target", default=None)
parser.add_argument("--sample-index", type=int, default=0)
args = parser.parse_args()
config = load_config(args.config)
output_dir = resolve_path(config["paths"]["result_dir"], args.config) / "output"
prediction = np.load(args.prediction or output_dir / "prediction.npy")
target = np.load(args.target or output_dir / "target.npy")
if prediction.ndim != 6 or target.ndim != 5 or prediction.shape[0] != target.shape[0]:
raise ValueError(f"unexpected prediction/target shapes: {prediction.shape}, {target.shape}")
ensemble_mean = prediction.mean(axis=1)
ensemble_std = prediction.std(axis=1)
rmse = float(np.sqrt(np.mean((ensemble_mean - target) ** 2)))
acc = _acc(prediction, target)
crps = _crps(prediction, target)
result_dir = resolve_path(config["paths"]["result_dir"], args.config)
result_dir.mkdir(parents=True, exist_ok=True)
np.save(result_dir / "rmse.npy", np.asarray(rmse, dtype=np.float32))
np.save(result_dir / "acc.npy", np.asarray(acc, dtype=np.float32))
np.save(result_dir / "crps.npy", np.asarray(crps, dtype=np.float32))
sample = min(max(args.sample_index, 0), len(target) - 1)
variable = config["validation"]["plot_variable"]
channel = config["data"]["variables"].index(variable)
truth = target[sample, channel].mean(axis=0)
mean = ensemble_mean[sample, channel].mean(axis=0)
spread = ensemble_std[sample, channel].mean(axis=0)
figure, axes = plt.subplots(1, 3, figsize=(12, 4), constrained_layout=True)
for axis, field, title in zip(axes, (truth, mean, spread), ("Target", "Ensemble mean", "Ensemble spread")):
image = axis.imshow(field, cmap="coolwarm")
axis.set_title(title)
axis.set_xticks([])
axis.set_yticks([])
figure.colorbar(image, ax=axis, fraction=0.046, pad=0.04)
figure.savefig(result_dir / "forecast.png", dpi=150)
loss_dir = resolve_path(config["paths"]["checkpoint_dir"], args.config)
if (loss_dir / "train_loss.npy").exists() and (loss_dir / "val_loss.npy").exists():
figure, axis = plt.subplots(figsize=(6, 4))
axis.plot(np.load(loss_dir / "train_loss.npy"), label="train")
axis.plot(np.linspace(0, len(np.load(loss_dir / "train_loss.npy")), len(np.load(loss_dir / "val_loss.npy"))), np.load(loss_dir / "val_loss.npy"), label="validation")
axis.set_xlabel("step")
axis.set_ylabel("loss")
axis.legend()
figure.tight_layout()
figure.savefig(result_dir / "loss.png", dpi=150)
print(f"RMSE={rmse:.6f} ACC={acc:.6f} CRPS={crps:.6f}")
print(f"saved results to {result_dir}")
if __name__ == "__main__":
main()