CRAI-ClimateExtremes / scripts /inference.py
zhangrenchao's picture
Publish CRAI-ClimateExtremes reproduction
b20ca9c verified
Raw
History Blame Contribute Delete
2.56 kB
"""Run bounded ensemble reconstruction from a unified checkpoint."""
from pathlib import Path
import argparse
import json
import sys
import numpy as np
import torch
import yaml
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.crai_climateextremes import CRAIClimateExtremes
def main():
parser = argparse.ArgumentParser()
parser.add_argument("--config", type=Path, default=ROOT / "conf/config.yaml")
args = parser.parse_args()
config_path = args.config if args.config.is_absolute() else ROOT / args.config
with open(config_path, encoding="utf-8") as handle:
cfg = yaml.safe_load(handle)
data = np.load(ROOT / cfg["data_path"])
inputs = torch.from_numpy(np.concatenate((data["observed"], data["valid_mask"]), axis=1))
checkpoint_path = ROOT / cfg["checkpoint_path"]
if not checkpoint_path.is_file():
raise FileNotFoundError(f"checkpoint not found: {checkpoint_path}; run scripts/train.py first")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=True)
if checkpoint.get("format_version") != "1.0" or not isinstance(checkpoint.get("model"), list):
raise ValueError(f"unsupported checkpoint format: {checkpoint_path}")
model_config = checkpoint.get("model_config", {})
predictions = []
for state in checkpoint["model"]:
model = CRAIClimateExtremes(**model_config).to(device)
model.load_state_dict(state); model.eval()
with torch.no_grad():
predictions.append(model(inputs.to(device)).cpu().numpy())
if not predictions:
raise ValueError(f"checkpoint contains no ensemble members: {checkpoint_path}")
members = np.stack(predictions)
output = ROOT / cfg["output_dir"]
output.mkdir(parents=True, exist_ok=True)
np.savez_compressed(
output / "predictions.npz", prediction=members.mean(0),
ensemble_std=members.std(0), target=data["target"],
observed=data["observed"], valid_mask=data["valid_mask"],
europe_mask=data["europe_mask"], index_ids=data["index_ids"],
index_names=data["index_names"],
)
metadata = {"ensemble_members": len(predictions), "checkpoint_semantics": "member state list in one checkpoint", "output_range": [0, 100]}
(output / "metadata.json").write_text(json.dumps(metadata, indent=2) + "\n")
print(f"predicted {inputs.shape[0]} samples with {len(predictions)} members")
if __name__ == "__main__":
main()