File size: 2,704 Bytes
20cdc88 | 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 | """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()
|