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