StableNN-Phys / scripts /result.py
zhangrenchao's picture
Upload folder using huggingface_hub
7f71cfd verified
Raw
History Blame Contribute Delete
3.36 kB
import json
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
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())
data = np.load(ROOT / cfg["paths"]["inference"])
pred, target = data["prediction"], data["target"]
if pred.shape != target.shape or pred.shape[1:] != (65, 68):
raise ValueError(f"Expected prediction and target shaped [N,65,68], got {pred.shape} and {target.shape}")
mass = np.concatenate((data["layer_mass"], data["layer_mass"]))
mass = mass / mass.mean()
error = pred - target
mad = (np.abs(error) * mass[None, None]).mean(axis=(0, 2))
bias = error.mean(axis=(0, 2))
flat_target = target[:, 1:].astype(np.float64).reshape(target.shape[0], -1)
flat_pred = pred[:, 1:].astype(np.float64).reshape(pred.shape[0], -1)
source_r2 = []
for truth, estimate in zip(flat_target, flat_pred):
denom = np.sum((truth - truth.mean()) ** 2)
source_r2.append(float(1.0 - np.sum((truth - estimate) ** 2) / max(denom, 1e-12)))
dt = float(data["dt_seconds"])
qt_pred, qt_target = pred[..., 34:], target[..., 34:]
column_water_pred = (qt_pred * data["layer_mass"][None, None]).sum(-1)
column_water_target = (qt_target * data["layer_mass"][None, None]).sum(-1)
lhf_evap = data["surface"][:, :, 1] / 2.5e6
adv_q = (0.5 * (data["advection"][:, :-1, 34:] + data["advection"][:, 1:, 34:]) *
data["layer_mass"][None, None]).sum(-1)
storage = np.diff(column_water_pred, axis=1) / dt
precipitation = np.maximum(0.0, lhf_evap + adv_q - storage)
residual = storage + precipitation - lhf_evap - adv_q
metrics = {"mass_weighted_mad_per_step": mad.tolist(), "bias_per_step": bias.tolist(),
"source_r2": source_r2, "mean_source_r2": float(np.mean(source_r2)),
"water_budget": {"mean_precipitation_kg_m2_s": float(precipitation.mean()),
"mean_abs_residual_kg_m2_s": float(np.abs(residual).mean())},
"per_step_length": int(len(mad)), "step_hours": cfg["data"]["step_hours"]}
out = ROOT / cfg["paths"]["evaluation"]
out.parent.mkdir(parents=True, exist_ok=True)
out.write_text(json.dumps(metrics, indent=2))
hours = data["lead_hours"]
fig, axes = plt.subplots(2, 1, figsize=(9, 7), constrained_layout=True)
axes[0].plot(hours, target[:, :, :34].mean((0, 2)), label="target sL")
axes[0].plot(hours, pred[:, :, :34].mean((0, 2)), "--", label="predicted sL")
ax2 = axes[0].twinx()
ax2.plot(hours, target[:, :, 34:].mean((0, 2)), color="tab:green", label="target qT")
ax2.plot(hours, pred[:, :, 34:].mean((0, 2)), "--", color="tab:red", label="predicted qT")
axes[0].set(xlabel="lead time (h)", title="SCM state rollout")
axes[0].legend(loc="upper left")
ax2.legend(loc="upper right")
axes[1].plot(hours[1:], precipitation.mean(0) * 86400.0, color="navy")
axes[1].set(xlabel="lead time (h)", ylabel="mm day-1", title="Diagnosed precipitation")
fig.savefig(out.parent / "state_precipitation_timeseries.png", dpi=150)
plt.close(fig)
print(f"saved {out}: per_step={len(mad)}, source_R2={metrics['mean_source_r2']:.4f}")
if __name__ == "__main__":
main()