Download scripts/inference.py from OneScience-Group/CRAI-ClimateExtremes: direct link, hf CLI and curl.
- Browser
- Download file 2.56 kB
-
https://huggingface.co/OneScience-Group/CRAI-ClimateExtremes/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/CRAI-ClimateExtremes/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/CRAI-ClimateExtremes/resolve/main/scripts/inference.py
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() | |