CoolFace
Modelpublic

OneScience-Group/CRAI-ClimateExtremes

sourceHugging Faceapache-2.0updated 12d agoView on Hugging Face
0likes24downloads
result.py82 linesDownload Raw Back to scripts
1"""Evaluate missing-region reconstruction and create a comparison figure."""2 3from pathlib import Path4import argparse5import json6import numpy as np7import torch8import torch.nn.functional as F9import yaml10 11 12def correlation(a, b):13    if a.size < 2 or np.std(a) == 0 or np.std(b) == 0:14        return 0.015    return float(np.corrcoef(a, b)[0, 1])16 17 18def rankdata(values):19    order = np.argsort(values, kind="mergesort")20    ranks = np.empty(len(values), dtype=np.float64)21    ranks[order] = np.arange(len(values), dtype=np.float64)22    unique, inverse, counts = np.unique(values, return_inverse=True, return_counts=True)23    del unique24    for group, count in enumerate(counts):25        if count > 1:26            positions = np.flatnonzero(inverse == group)27            ranks[positions] = ranks[positions].mean()28    return ranks29 30 31def main():32    parser = argparse.ArgumentParser()33    root = Path(__file__).resolve().parents[1]34    parser.add_argument("--config", type=Path, default=root / "conf/config.yaml")35    args = parser.parse_args()36    config_path = args.config if args.config.is_absolute() else root / args.config37    with open(config_path, encoding="utf-8") as handle:38        cfg = yaml.safe_load(handle)39    archive = np.load(root / cfg["output_dir"] / "predictions.npz")40    pred, target = archive["prediction"], archive["target"]41    missing = archive["europe_mask"][None, None] * (1 - archive["valid_mask"])42    selected = missing.astype(bool)43    error = pred[selected] - target[selected]44    sample_spearman = []45    for i in range(len(pred)):46        mask = selected[i, 0]47        sample_spearman.append(correlation(rankdata(pred[i, 0][mask]), rankdata(target[i, 0][mask])))48    kernel = torch.ones(1, 1, 3, 3) / 849    kernel[0, 0, 1, 1] = 050    pred_neighbor = F.conv2d(torch.from_numpy(pred), kernel, padding=1).numpy()51    target_neighbor = F.conv2d(torch.from_numpy(target), kernel, padding=1).numpy()52    metrics = {53        "missing_rmse": float(np.sqrt(np.mean(error ** 2))),54        "sample_spearman_mean": float(np.mean(sample_spearman)),55        "sample_spearman": [float(x) for x in sample_spearman],56        "missing_bias": float(np.mean(error)),57        "neighborhood_spatial_correlation_prediction": correlation(pred[selected], pred_neighbor[selected]),58        "neighborhood_spatial_correlation_target": correlation(target[selected], target_neighbor[selected]),59        "missing_points": int(selected.sum()),60    }61    output = root / cfg["evaluation_dir"]62    output.mkdir(parents=True, exist_ok=True)63    (output / "metrics.json").write_text(json.dumps(metrics, indent=2) + "\n")64    import matplotlib65    matplotlib.use("Agg")66    import matplotlib.pyplot as plt67    fig, axes = plt.subplots(1, 3, figsize=(13, 4), constrained_layout=True)68    sample = 069    fields = [archive["observed"][sample, 0], target[sample, 0], pred[sample, 0]]70    titles = ["Irregular observations", "Synthetic truth", "CRAI reconstruction"]71    for axis, field, title in zip(axes, fields, titles):72        image = axis.imshow(np.where(archive["europe_mask"] > 0, field, np.nan), origin="lower", vmin=0, vmax=100, cmap="RdYlBu_r")73        axis.set_title(title); axis.set_axis_off()74    fig.colorbar(image, ax=axes, label="Extreme index (%)", shrink=0.8)75    fig.suptitle(f"Missing RMSE={metrics['missing_rmse']:.3f} | Spearman={metrics['sample_spearman_mean']:.3f} | Bias={metrics['missing_bias']:.3f}")76    fig.savefig(output / "comparison.png", dpi=160); plt.close(fig)77    print(json.dumps(metrics))78 79 80if __name__ == "__main__":81    main()82