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