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