Download scripts/result.py from OneScience-Group/CorrDiff: direct link, hf CLI and curl.
- Browser
- Download file 3.15 kB
-
https://huggingface.co/OneScience-Group/CorrDiff/resolve/main/scripts/result.py
- Command line
-
hf download hf://OneScience-Group/CorrDiff/scripts/result.py
-
curl -L -o result.py https://huggingface.co/OneScience-Group/CorrDiff/resolve/main/scripts/result.py
3.15 kB
| """Compute deterministic and probabilistic CorrDiff metrics and plots.""" | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import matplotlib | |
| matplotlib.use("Agg") | |
| import matplotlib.pyplot as plt | |
| import numpy as np | |
| import yaml | |
| ROOT = Path(__file__).resolve().parents[1] | |
| def crps_ensemble(ensemble, target): | |
| first = np.abs(ensemble - target[None]).mean(0) | |
| sorted_members = np.sort(ensemble, axis=0) | |
| m = ensemble.shape[0] | |
| weights = (2 * np.arange(1, m + 1) - m - 1).reshape(m, 1, 1, 1, 1) | |
| return first - (sorted_members * weights).sum(0) / m**2 | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--config", default=str(ROOT / "conf/config.yaml")) | |
| args = parser.parse_args() | |
| config = yaml.safe_load(Path(args.config).read_text(encoding="utf-8")) | |
| archive = np.load(ROOT / config["paths"]["predictions"]) | |
| if "protocol" not in archive or archive["protocol"].ndim != 0 or str(archive["protocol"].item()) != config["data"]["protocol"]: | |
| raise ValueError("Prediction protocol does not match the configured protocol") | |
| if "data_source" not in archive or archive["data_source"].ndim != 0 or not str(archive["data_source"].item()): | |
| raise ValueError("Prediction data_source must be a non-empty scalar") | |
| protocol = str(archive["protocol"].item()) | |
| data_source = str(archive["data_source"].item()) | |
| ensemble, target = archive["ensemble"], archive["target"] | |
| mean, spread = ensemble.mean(0), ensemble.std(0) | |
| axes = (0, 2, 3) | |
| mae = np.abs(mean - target).mean(axis=axes) | |
| rmse = np.sqrt(((mean - target) ** 2).mean(axis=axes)) | |
| crps = crps_ensemble(ensemble, target).mean(axis=axes) | |
| spread_value = spread.mean(axis=axes) | |
| names = config["data"]["target_variables"] | |
| metrics = {name: {"mae": float(mae[i]), "rmse": float(rmse[i]), "crps": float(crps[i]), | |
| "ensemble_spread": float(spread_value[i])} for i, name in enumerate(names)} | |
| metrics["aggregate"] = {key: float(np.mean([metrics[n][key] for n in names])) | |
| for key in ("mae", "rmse", "crps", "ensemble_spread")} | |
| output = ROOT / config["paths"]["evaluation_dir"] | |
| output.mkdir(parents=True, exist_ok=True) | |
| payload = {"metrics": metrics, "protocol": protocol, "data_source": data_source} | |
| (output / "metrics.json").write_text(json.dumps(payload, indent=2) + "\n") | |
| figure, plot_axes = plt.subplots(len(names), 4, figsize=(13, 3 * len(names))) | |
| for channel, name in enumerate(names): | |
| fields = (target[0, channel], mean[0, channel], spread[0, channel], mean[0, channel] - target[0, channel]) | |
| titles = ("target", "ensemble mean", "ensemble spread", "mean error") | |
| for axis, field, title in zip(plot_axes[channel], fields, titles): | |
| axis.imshow(field, cmap="coolwarm" if title == "mean error" else "viridis") | |
| axis.set_title(f"{name}: {title}"); axis.axis("off") | |
| figure.tight_layout(); figure.savefig(output / "ensemble_diagnostics.png", dpi=120); plt.close(figure) | |
| print(json.dumps(payload, indent=2)); print(f"evaluation={output}") | |
| if __name__ == "__main__": | |
| main() | |