| """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() |
|
|