EchoCast-3D / scripts /result.py
zhangrenchao's picture
Publish EchoCast-3D reproduction
e0a6aa0 verified
Raw
History Blame Contribute Delete
2.18 kB
"""Evaluate the ensemble and render a 2D composite-reflectivity summary."""
from __future__ import annotations
import argparse
import json
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
import sys
ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
from model.echocast_3d import load_config, evaluate_ensemble
def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--config", default="conf/config.yaml")
parser.add_argument("--predictions", default="result/output/predictions.npz")
parser.add_argument("--metrics", default="result/evaluation/metrics.json")
parser.add_argument("--figure", default="result/evaluation/comparison.png")
args = parser.parse_args()
config = load_config(ROOT / args.config)
with np.load(ROOT / args.predictions) as data:
ensemble = data["ensemble"]
truth = data["truth"]
valid = data["validity"].astype(bool)
result = evaluate_ensemble(ensemble, truth, valid, config["evaluation"]["thresholds_dbz"])
metrics_path = ROOT / args.metrics
metrics_path.parent.mkdir(parents=True, exist_ok=True)
metrics_path.write_text(json.dumps(result, indent=2), encoding="utf-8")
prediction = ensemble.mean(axis=0)
fig, axes = plt.subplots(2, 5, figsize=(15, 6), constrained_layout=True)
for lead in range(5):
for row, field in enumerate((truth, prediction)):
composite = np.where(valid[lead], field[lead], np.nan).max(axis=0)
image = axes[row, lead].imshow(composite.T, origin="lower", vmin=0, vmax=60, cmap="turbo", aspect="auto")
axes[row, lead].set_title(f"{'Truth' if row == 0 else 'Ensemble mean'} +{(lead + 1) * 6} min")
axes[row, lead].set_xlabel("azimuth index")
if lead == 0:
axes[row, lead].set_ylabel("range bin")
fig.colorbar(image, ax=axes, label="composite reflectivity (dBZ)", shrink=0.8)
figure_path=ROOT/args.figure;figure_path.parent.mkdir(parents=True,exist_ok=True);fig.savefig(figure_path,dpi=120)
plt.close(fig)
print(json.dumps(result, indent=2))
if __name__ == "__main__":
main()