OneScience-Group/FuXi-DA
026
1"""Tile diagnostics for analysis, forecasts, variable groups, and localization."""2import argparse3import json4from pathlib import Path5 6import matplotlib.pyplot as plt7import numpy as np8import torch9import sys, yaml10ROOT = Path(__file__).resolve().parents[1]11sys.path.insert(0, str(ROOT))12 13from model.fuxi_da import CompactForecastProxy, FuXiDA, make_sample14 15GROUPS = {"Z": slice(0, 13), "T": slice(13, 26), "U": slice(26, 39), "V": slice(39, 52), "R": slice(52, 65), "surface": slice(65, 70)}16 17 18def weighted_rmse(prediction, target, latitude):19 weight = torch.cos(torch.deg2rad(latitude)).clamp_min(0)20 weight = weight / weight.sum()21 return torch.sqrt(((prediction - target).square() * weight[None, :, None]).sum(dim=(-2, -1)) / prediction.shape[-1]).mean().item()22 23 24def main():25 parser = argparse.ArgumentParser()26 parser.add_argument("--config", default="conf/config.yaml")27 parser.add_argument("--checkpoint")28 parser.add_argument("--output")29 parser.add_argument("--metrics")30 args = parser.parse_args()31 cfg = yaml.safe_load((ROOT / args.config).read_text())32 model = FuXiDA(cfg["model"]["base_channels"])33 checkpoint_path = Path(args.checkpoint) if args.checkpoint else ROOT / cfg["paths"]["checkpoint"]34 if checkpoint_path.exists():35 model.load_state_dict(torch.load(checkpoint_path, map_location="cpu", weights_only=True)["model"])36 model.eval()37 proxy = CompactForecastProxy()38 39 rows = []40 with torch.no_grad():41 for tile_id in cfg["data"]["tile_ids"]:42 sample = make_sample(tile_id, 20, cfg["data"]["tile_size"], cfg["data"]["missing_probability"])43 analysis = model(sample["background"][None], sample["obs"][None])[0]44 row = {"tile_id": tile_id, "origin": sample["origin"].tolist(), "background": weighted_rmse(sample["background"], sample["target"], sample["latitude"]), "correction": weighted_rmse(sample["correction"], sample["target"], sample["latitude"]), "analysis": weighted_rmse(analysis, sample["target"], sample["latitude"])}45 row["variable_groups"] = {name: weighted_rmse(analysis[index], sample["target"][index], sample["latitude"]) for name, index in GROUPS.items()}46 state = analysis47 row["forecast_steps"] = []48 for lead in range(cfg["train"]["forecast_steps"]):49 state = proxy(state[None])[0]50 row["forecast_steps"].append(weighted_rmse(state, sample["forecast_targets"][lead], sample["latitude"]))51 rows.append(row)52 53 sample = make_sample(cfg["data"]["tile_ids"][0], 21, cfg["data"]["tile_size"], 0.0)54 base = model(sample["background"][None], sample["obs"][None])[0]55 perturbed_obs = sample["obs"].clone()56 center = cfg["data"]["tile_size"] // 257 perturbed_obs[9 - 8, center, center] += 1.058 response = (model(sample["background"][None], perturbed_obs[None])[0] - base).square().sum(0)59 yy, xx = torch.meshgrid(torch.arange(cfg["data"]["tile_size"]), torch.arange(cfg["data"]["tile_size"]), indexing="ij")60 radius = torch.sqrt((yy - center).float().square() + (xx - center).float().square())61 total_energy = response.sum().clamp_min(1e-12)62 localization = {"perturbation": "AGRI channel 9 +1 K at tile center", "energy_weighted_radius_gridpoints": (response * radius).sum().div(total_energy).item(), "energy_within_radius_4": response[radius <= 4].sum().div(total_energy).item()}63 64 metrics = {"coverage_complete": False, "claim": "diagnostic metrics on selected aligned tiles; not complete global scores", "tiles": rows, "increment_localization": localization}65 metrics_path = Path(args.metrics) if args.metrics else ROOT / cfg["paths"]["evaluation"]66 metrics_path.parent.mkdir(parents=True, exist_ok=True); metrics_path.write_text(json.dumps(metrics, indent=2) + "\n")67 labels = ["background", "correction", "analysis"]68 values = [np.mean([row[label] for row in rows]) for label in labels]69 fig, axes = plt.subplots(1, 2, figsize=(10, 4))70 axes[0].bar(labels, values, color=["#596780", "#d69b36", "#287f71"])71 axes[0].set_ylabel("Latitude-weighted RMSE")72 axes[0].set_title("Selected aligned tiles (incomplete coverage)")73 group_values = [np.mean([row["variable_groups"][name] for row in rows]) for name in GROUPS]74 axes[1].bar(list(GROUPS), group_values, color="#287f71")75 axes[1].set_title("Analysis RMSE by variable group")76 fig.suptitle("FuXi-DA procedural protocol diagnostics")77 fig.tight_layout()78 output_path = Path(args.output) if args.output else ROOT / cfg["paths"]["figure"]79 output_path.parent.mkdir(parents=True, exist_ok=True); fig.savefig(output_path, dpi=160)80 plt.close(fig)81 print(json.dumps(metrics, indent=2))82 83 84if __name__ == "__main__":85 main()86 