"""Compute per-lead latitude-weighted model and persistence RMSE/MBE.""" import json from pathlib import Path import sys import matplotlib.pyplot as plt import numpy as np import yaml ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT)) from model.fuxi_ocean import ocean_channel_depth_indices, weighted_error_sums def aggregate(prediction, truth, latitude, depth_mask): channel_depth = ocean_channel_depth_indices() mask = np.take(depth_mask, channel_depth, axis=1) squared, bias, weight = weighted_error_sums(prediction, truth, latitude, mask) return np.sqrt(squared / np.maximum(weight, 1)), bias / np.maximum(weight, 1) def main(): config = yaml.safe_load((ROOT / "conf/config.yaml").read_text()) data = np.load(ROOT / config["paths"]["inference"]) model_rmse, model_mbe, persistence_rmse, persistence_mbe = [], [], [], [] for lead in range(data["prediction"].shape[1]): rmse, mbe = aggregate(data["prediction"][:, lead], data["truth"][:, lead], data["latitude_deg"], data["depth_mask"]) prmse, pmbe = aggregate(data["initial"], data["truth"][:, lead], data["latitude_deg"], data["depth_mask"]) model_rmse.append(rmse); model_mbe.append(mbe); persistence_rmse.append(prmse); persistence_mbe.append(pmbe) arrays = [np.asarray(value) for value in (model_rmse, model_mbe, persistence_rmse, persistence_mbe)] lead_hours = data["lead_hours"].tolist() groups = json.loads(str(data["variable_groups"])) def grouped(values, i): return {name: {"unit": spec["unit"], "rmse": float(np.mean(values[0][i, slice(*spec["channels"])])), "mbe": float(np.mean(values[1][i, slice(*spec["channels"])]))} for name, spec in groups.items()} metrics = {"output_kind": str(data["output_kind"]), "format_version": str(data["format_version"]), "checkpoint_source": str(data["checkpoint_source"]), "sample_count": int(data["sample_count"]), "coverage_fraction": float(data["coverage_fraction"]), "is_complete_global": bool(data["is_complete_global"]), "synthetic": bool(data["synthetic"]), "aggregation": "latitude-weighted sampled tiles; never across unit groups", "output_shape": data["output_shape"].tolist(), "variable_groups": groups, "per_lead": [{"lead_hours": hour, "model_by_group": grouped(arrays[:2], i), "persistence_by_group": grouped(arrays[2:], i), "model_by_channel": {"rmse": arrays[0][i].tolist(), "mbe": arrays[1][i].tolist()}, "persistence_by_channel": {"rmse": arrays[2][i].tolist(), "mbe": arrays[3][i].tolist()}} for i, hour in enumerate(lead_hours)]} path = ROOT / config["paths"]["evaluation_metrics"]; path.parent.mkdir(parents=True, exist_ok=True) path.write_text(json.dumps(metrics, indent=2) + "\n") figure, axes = plt.subplots(1, 5, figsize=(18, 3.5)) for axis, (name, spec) in zip(axes, groups.items()): channel_slice = slice(*spec["channels"]) axis.plot(lead_hours, arrays[0][:, channel_slice].mean(1), "o-", label="model") axis.plot(lead_hours, arrays[2][:, channel_slice].mean(1), "s--", label="persistence") axis.set(title=name, xlabel="Lead (h)", ylabel=f"RMSE ({spec['unit']})"); axis.legend() figure.tight_layout() plot = ROOT / config["paths"]["evaluation_plot"]; figure.savefig(plot, dpi=160); plt.close(figure) print(f"metrics={path.relative_to(ROOT)} plot={plot.relative_to(ROOT)} leads={lead_hours}") if __name__ == "__main__": main()