File size: 4,812 Bytes
15aff58
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
"""Tile diagnostics for analysis, forecasts, variable groups, and localization."""
import argparse
import json
from pathlib import Path

import matplotlib.pyplot as plt
import numpy as np
import torch
import sys, yaml
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))

from model.fuxi_da import CompactForecastProxy, FuXiDA, make_sample

GROUPS = {"Z": slice(0, 13), "T": slice(13, 26), "U": slice(26, 39), "V": slice(39, 52), "R": slice(52, 65), "surface": slice(65, 70)}


def weighted_rmse(prediction, target, latitude):
    weight = torch.cos(torch.deg2rad(latitude)).clamp_min(0)
    weight = weight / weight.sum()
    return torch.sqrt(((prediction - target).square() * weight[None, :, None]).sum(dim=(-2, -1)) / prediction.shape[-1]).mean().item()


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", default="conf/config.yaml")
    parser.add_argument("--checkpoint")
    parser.add_argument("--output")
    parser.add_argument("--metrics")
    args = parser.parse_args()
    cfg = yaml.safe_load((ROOT / args.config).read_text())
    model = FuXiDA(cfg["model"]["base_channels"])
    checkpoint_path = Path(args.checkpoint) if args.checkpoint else ROOT / cfg["paths"]["checkpoint"]
    if checkpoint_path.exists():
        model.load_state_dict(torch.load(checkpoint_path, map_location="cpu", weights_only=True)["model"])
    model.eval()
    proxy = CompactForecastProxy()

    rows = []
    with torch.no_grad():
        for tile_id in cfg["data"]["tile_ids"]:
            sample = make_sample(tile_id, 20, cfg["data"]["tile_size"], cfg["data"]["missing_probability"])
            analysis = model(sample["background"][None], sample["obs"][None])[0]
            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"])}
            row["variable_groups"] = {name: weighted_rmse(analysis[index], sample["target"][index], sample["latitude"]) for name, index in GROUPS.items()}
            state = analysis
            row["forecast_steps"] = []
            for lead in range(cfg["train"]["forecast_steps"]):
                state = proxy(state[None])[0]
                row["forecast_steps"].append(weighted_rmse(state, sample["forecast_targets"][lead], sample["latitude"]))
            rows.append(row)

        sample = make_sample(cfg["data"]["tile_ids"][0], 21, cfg["data"]["tile_size"], 0.0)
        base = model(sample["background"][None], sample["obs"][None])[0]
        perturbed_obs = sample["obs"].clone()
        center = cfg["data"]["tile_size"] // 2
        perturbed_obs[9 - 8, center, center] += 1.0
        response = (model(sample["background"][None], perturbed_obs[None])[0] - base).square().sum(0)
        yy, xx = torch.meshgrid(torch.arange(cfg["data"]["tile_size"]), torch.arange(cfg["data"]["tile_size"]), indexing="ij")
        radius = torch.sqrt((yy - center).float().square() + (xx - center).float().square())
        total_energy = response.sum().clamp_min(1e-12)
        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()}

    metrics = {"coverage_complete": False, "claim": "diagnostic metrics on selected aligned tiles; not complete global scores", "tiles": rows, "increment_localization": localization}
    metrics_path = Path(args.metrics) if args.metrics else ROOT / cfg["paths"]["evaluation"]
    metrics_path.parent.mkdir(parents=True, exist_ok=True); metrics_path.write_text(json.dumps(metrics, indent=2) + "\n")
    labels = ["background", "correction", "analysis"]
    values = [np.mean([row[label] for row in rows]) for label in labels]
    fig, axes = plt.subplots(1, 2, figsize=(10, 4))
    axes[0].bar(labels, values, color=["#596780", "#d69b36", "#287f71"])
    axes[0].set_ylabel("Latitude-weighted RMSE")
    axes[0].set_title("Selected aligned tiles (incomplete coverage)")
    group_values = [np.mean([row["variable_groups"][name] for row in rows]) for name in GROUPS]
    axes[1].bar(list(GROUPS), group_values, color="#287f71")
    axes[1].set_title("Analysis RMSE by variable group")
    fig.suptitle("FuXi-DA procedural protocol diagnostics")
    fig.tight_layout()
    output_path = Path(args.output) if args.output else ROOT / cfg["paths"]["figure"]
    output_path.parent.mkdir(parents=True, exist_ok=True); fig.savefig(output_path, dpi=160)
    plt.close(fig)
    print(json.dumps(metrics, indent=2))


if __name__ == "__main__":
    main()