Team Ai
Modelpublic

OneScience-Group/FuXi-DA

sourceHugging Faceapache-2.0updated 25d agoView on Hugging Face
0likes26downloads
result.py86 linesDownload Raw Back to scripts
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