zhangrenchao's picture
Publish CausalModelEvaluation engineering reproduction
20cdc88 verified
Raw
History Blame Contribute Delete
2.69 kB
"""Fit statistical CME networks and the precipitation-constraint GP."""
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 (FORMAT_VERSION, LaggedPartialCorrelationCME,
PrecipitationConstraintGP, asymmetric_f1)
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
np.random.seed(int(config["seed"]))
data = np.load(ROOT / config["paths"]["dataset"])
if str(data["format_version"]) != config["data"]["format_version"]:
raise ValueError("incompatible synthetic data format")
estimator = LaggedPartialCorrelationCME(**config["model"])
references = [estimator.fit(series) for series in data["reference_series"]]
model_networks, scores = [], []
for model_series in data["model_series"]:
networks = [estimator.fit(series) for series in model_series]
model_networks.append(networks)
scores.append(float(np.mean([asymmetric_f1(ref, net, int(config["evaluation"]["lag_tolerance"]))["f1"]
for ref, net in zip(references, networks)])))
gp = PrecipitationConstraintGP(int(config["seed"])).fit(np.asarray(scores), data["delta_precipitation"])
checkpoint = {
"version": FORMAT_VERSION, "config": config, "network_config": estimator.config(),
"reference_networks": [network.state_dict() for network in references],
"model_networks": [[network.state_dict() for network in networks] for networks in model_networks],
"model_f1": torch.tensor(scores), "gp": gp.state_dict(),
"metadata": {"method": "lagged target-history conditional regression ParCorr approximation",
"gradient_training": False, "tensor_layout": "source,target,lag", "time_step_days": 3}
}
checkpoint_path = ROOT / config["paths"]["checkpoint"]
metrics_path = ROOT / config["paths"]["training_metrics"]
checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
metrics_path.parent.mkdir(parents=True, exist_ok=True)
torch.save(checkpoint, checkpoint_path)
metrics_path.write_text(json.dumps({"model_f1": scores, "gp_kernel": str(gp.model.kernel_),
"paper_alpha": config["paper_model"]["alpha"],
"engineering_alpha": config["model"]["alpha"]}, indent=2) + "\n")
print(f"checkpoint={checkpoint_path.relative_to(ROOT)} networks={len(references) * (1 + len(model_networks))} gp_samples={len(scores)}")
if __name__ == "__main__":
main()