zhangrenchao's picture
Publish CausalModelEvaluation engineering reproduction
20cdc88 verified
Raw
History Blame Contribute Delete
2.7 kB
"""Reload the CME checkpoint and export complete network and projection outputs."""
import json
import sys
from pathlib import Path
import numpy as np
import torch
import yaml
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.causalmodelevaluation import CausalNetwork, FORMAT_VERSION, PrecipitationConstraintGP
def stack_networks(states, field):
return np.stack([[getattr(CausalNetwork.from_state_dict(state), field).numpy() for state in group] for group in states])
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
checkpoint = torch.load(ROOT / config["paths"]["checkpoint"], map_location="cpu", weights_only=False)
if checkpoint["version"] != FORMAT_VERSION:
raise ValueError("unsupported checkpoint version")
reference_states = [checkpoint["reference_networks"]]
model_states = checkpoint["model_networks"]
f1 = checkpoint["model_f1"].numpy()
gp = PrecipitationConstraintGP.from_state_dict(checkpoint["gp"], int(config["seed"]))
query = np.sort(np.unique(np.append(f1, float(config["evaluation"]["reference_f1_for_projection"]))))
mean, lower, upper = gp.predict(query)
source = np.load(ROOT / config["paths"]["dataset"])
output = ROOT / config["paths"]["inference"]
output.parent.mkdir(parents=True, exist_ok=True)
np.savez_compressed(
output, edges=stack_networks(model_states, "edges"), pvalues=stack_networks(model_states, "pvalues"),
mci=stack_networks(model_states, "mci"), reference_edges=stack_networks(reference_states, "edges")[0],
reference_pvalues=stack_networks(reference_states, "pvalues")[0],
reference_mci=stack_networks(reference_states, "mci")[0], model_f1=f1,
reference_precipitation=source["reference_precipitation"], model_precipitation=source["model_precipitation"],
delta_precipitation=source["delta_precipitation"], latitude_degrees=source["latitude_degrees"],
longitude_degrees=source["longitude_degrees"], gp_query_f1=query, gp_mean_delta_precipitation=mean,
gp_lower_95=lower, gp_upper_95=upper, seasons=source["seasons"],
network_metadata=np.asarray(json.dumps(checkpoint["metadata"])),
projection_metadata=np.asarray(json.dumps({"kernel": str(gp.model.kernel_), "confidence": 0.95,
"input": "CME asymmetric F1", "output": "delta precipitation"})),
format_version=np.asarray(FORMAT_VERSION))
print(f"inference={output.relative_to(ROOT)} edges={stack_networks(model_states, 'edges').shape} reference={stack_networks(reference_states, 'edges')[0].shape}")
if __name__ == "__main__":
main()