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