zhangrenchao's picture
Publish CausalModelEvaluation engineering reproduction
20cdc88 verified
Raw
History Blame Contribute Delete
4.66 kB
"""Compute CME, precipitation metrics, GP intervals, and the task figure."""
import json
import sys
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
import yaml
from scipy import stats
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.causalmodelevaluation import CausalNetwork, asymmetric_f1, taylor_s_score
def network(edges, pvalues, mci):
return CausalNetwork(torch.from_numpy(edges), torch.from_numpy(pvalues), torch.from_numpy(mci))
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
data = np.load(ROOT / config["paths"]["inference"])
tolerance = int(config["evaluation"]["lag_tolerance"])
by_model, model_s = [], []
weights = np.cos(np.deg2rad(data["latitude_degrees"]))[:, None] * np.ones((1, len(data["longitude_degrees"])))
for model_index in range(data["edges"].shape[0]):
seasonal = []
for season in range(data["edges"].shape[1]):
seasonal.append(asymmetric_f1(
network(data["reference_edges"][season], data["reference_pvalues"][season], data["reference_mci"][season]),
network(data["edges"][model_index, season], data["pvalues"][model_index, season], data["mci"][model_index, season]), tolerance))
taylor = taylor_s_score(data["reference_precipitation"], data["model_precipitation"][model_index], weights)
model_s.append(taylor["s_score"])
edge_count = int(data["edges"][model_index].sum())
possible = int(np.prod(data["edges"][model_index].shape))
by_model.append({"model": model_index, "cme_f1": float(np.mean([x["f1"] for x in seasonal])),
"seasonal_cme": seasonal, "edge_count": edge_count, "edge_density": edge_count / possible,
"taylor": taylor, "delta_precipitation": float(data["delta_precipitation"][model_index])})
f1 = np.asarray([item["cme_f1"] for item in by_model])
s_scores = np.asarray(model_s)
delta = data["delta_precipitation"]
f1_delta = stats.pearsonr(f1, delta)
f1_s = stats.pearsonr(f1, s_scores)
report = {
"models": by_model,
"reference": {"edge_count": int(data["reference_edges"].sum()),
"edge_density": float(data["reference_edges"].mean())},
"correlations": {"f1_vs_delta_precipitation": {"r": float(f1_delta.statistic), "pvalue": float(f1_delta.pvalue)},
"f1_vs_taylor_s": {"r": float(f1_s.statistic), "pvalue": float(f1_s.pvalue)}},
"gp_projection": {"query_f1": data["gp_query_f1"].tolist(), "mean": data["gp_mean_delta_precipitation"].tolist(),
"lower_95": data["gp_lower_95"].tolist(), "upper_95": data["gp_upper_95"].tolist(),
"metadata": json.loads(str(data["projection_metadata"]))},
"protocol": {"direction_sign_required": True, "lag_tolerance_steps": tolerance,
"lag_tolerance_days": tolerance * int(config["data"]["time_step_days"]),
"paper_alpha": config["paper_model"]["alpha"], "engineering_alpha": config["model"]["alpha"]}
}
output = ROOT / config["paths"]["evaluation_dir"]
output.mkdir(parents=True, exist_ok=True)
(output / "metrics.json").write_text(json.dumps(report, indent=2) + "\n")
order = np.argsort(data["gp_query_f1"])
figure, axes = plt.subplots(1, 2, figsize=(11, 4.3))
axes[0].scatter(f1, s_scores, c=np.arange(len(f1)), cmap="viridis", s=55)
axes[0].set(xlabel="CME asymmetric F1", ylabel="Precipitation Taylor S-score", title=f"Network and precipitation skill (R={f1_s.statistic:.2f})")
axes[1].scatter(f1, delta, color="#9b3a2e", label="Synthetic models")
axes[1].plot(data["gp_query_f1"][order], data["gp_mean_delta_precipitation"][order], color="#173f5f", label="RBF + white GP")
axes[1].fill_between(data["gp_query_f1"][order], data["gp_lower_95"][order], data["gp_upper_95"][order], color="#4f8fba", alpha=0.25, label="95% interval")
axes[1].set(xlabel="CME asymmetric F1", ylabel="Delta precipitation", title="Constrained precipitation relationship")
axes[1].legend(fontsize=8)
figure.tight_layout()
figure.savefig(output / "cme_task.png", dpi=160)
plt.close(figure)
if not np.isfinite([f1_delta.statistic, f1_s.statistic, *data["gp_mean_delta_precipitation"]]).all():
raise FloatingPointError("evaluation contains non-finite values")
print(f"evaluation={output.relative_to(ROOT)} models={len(by_model)} figure=cme_task.png")
if __name__ == "__main__":
main()