"""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()