CESM-SeasonalML / scripts /inference.py
zhangrenchao's picture
Publish CESM-SeasonalML engineering reproduction
b300acf verified
Raw
History Blame Contribute Delete
4.66 kB
"""Generate complete test probabilities, classes, centroids, and years."""
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.cesm_seasonal_ml import (CHECKPOINT_FORMAT_VERSION, DATA_FORMAT_VERSION, FeedForwardNN,
SeasonalLSTM, StablePrecipKMeans)
def load_checkpoint(path):
try:
return torch.load(path, map_location="cpu", weights_only=False)
except TypeError:
return torch.load(path, map_location="cpu")
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
checkpoint = load_checkpoint(ROOT / config["paths"]["checkpoint"])
data = np.load(ROOT / config["data"]["path"])
required = {"format_version", "model", "model_config", "data_spec", "season"}
if not required.issubset(checkpoint):
raise ValueError(f"checkpoint missing required top-level keys: {sorted(required - checkpoint.keys())}")
if checkpoint["format_version"] != CHECKPOINT_FORMAT_VERSION:
raise ValueError("checkpoint format version mismatch")
spec, model_config = checkpoint["data_spec"], checkpoint["model_config"]
if str(data["format_version"]) != DATA_FORMAT_VERSION or spec["format_version"] != DATA_FORMAT_VERSION:
raise ValueError("data format version mismatch")
season = checkpoint["season"]
if season != spec["season"] or season not in [str(value) for value in data["seasons"]]:
raise ValueError("checkpoint/data season mismatch")
checks = (("rf_manifest", spec["rf_manifest"]), ("nn_manifest", spec["nn_manifest"]),
("class_names", spec["class_names"]))
for key, expected in checks:
if data[key].tolist() != expected:
raise ValueError(f"checkpoint/data {key} mismatch")
shapes = (("rf_features", "rf_shape"), ("nn_features", "nn_shape"),
("eof_sequence", "eof_shape"), ("precipitation", "precipitation_shape"))
for key, shape_key in shapes:
if list(data[key].shape) != spec[shape_key]:
raise ValueError(f"checkpoint/data {key} shape mismatch")
season_index = [str(x) for x in data["seasons"]].index(season)
test = data["split"] == 2
cluster = StablePrecipKMeans.from_state_dict(checkpoint["cluster"])
target = cluster.transform(data["precipitation"][season_index, test])
probabilities = {}
models, settings = checkpoint["model"], model_config["settings"]
if "rf" in models:
probabilities["rf"] = models["rf"].predict_proba(data["rf_features"][season_index, test])
if "xgboost" in models:
probabilities["xgboost"] = models["xgboost"].predict_proba(data["rf_features"][season_index, test])
if "nn" in models:
model = FeedForwardNN(input_size=model_config["nn_features"], hidden=tuple(settings["nn"]["hidden"]), dropout=settings["nn"]["dropout"], classes=model_config["classes"])
model.load_state_dict(models["nn"]); model.eval()
with torch.no_grad(): probabilities["nn"] = torch.softmax(model(torch.from_numpy(data["nn_features"][season_index, test]).float()), 1).numpy()
if "lstm" in models:
lstm = settings["lstm"]
model = SeasonalLSTM(input_size=model_config["eof_channels"], hidden_size=lstm["hidden_size"], dense_size=lstm["dense_size"], dropout=lstm["dropout"], classes=model_config["classes"])
model.load_state_dict(models["lstm"]); model.eval()
sequence = data["eof_sequence"][season_index, test, -model_config["history_length"]:]
with torch.no_grad(): probabilities["lstm"] = torch.softmax(model(torch.from_numpy(sequence).float()), 1).numpy()
primary = model_config["primary"] if model_config["primary"] in probabilities else next(iter(probabilities))
output = ROOT / config["paths"]["inference"]
output.parent.mkdir(parents=True, exist_ok=True)
payload = {"format_version": DATA_FORMAT_VERSION, "season": season, "years": data["years"][test], "target": target, "centroids": cluster.centroids_,
"latitude": data["latitude"], "longitude": data["longitude"], "primary_model": primary,
"class_names": data["class_names"], "target_precipitation": data["precipitation"][season_index, test]}
for name, probability in probabilities.items():
payload[f"probability_{name}"] = probability
payload[f"prediction_{name}"] = probability.argmax(axis=1)
np.savez_compressed(output, **payload)
print(f"predictions={output.relative_to(ROOT)} season={season} samples={test.sum()} primary={primary}")
if __name__ == "__main__":
main()