Download scripts/inference.py from OneScience-Group/Ai2_Climate_Emulator: direct link, hf CLI and curl.
- Browser
- Download file 6.2 kB
-
https://huggingface.co/OneScience-Group/Ai2_Climate_Emulator/resolve/main/scripts/inference.py
- Command line
-
hf download hf://OneScience-Group/Ai2_Climate_Emulator/scripts/inference.py
-
curl -L -o inference.py https://huggingface.co/OneScience-Group/Ai2_Climate_Emulator/resolve/main/scripts/inference.py
6.2 kB
| """Run ACE autoregressive rollout from a checkpoint.""" | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import yaml | |
| if __package__ in (None, ""): | |
| sys.path.insert(0, str(Path(__file__).resolve().parents[2])) | |
| from ACE.model.data import make_fake_pairs, save_fake_pairs | |
| from ACE.model.ace import ACEModel, ACEModelConfig | |
| from ACE.model.normalization import ACEDataNormalizer | |
| from ACE.model.paths import CHECKPOINT_PATH, GENERATED_DATA_PATH, INFER_DIR, configured_path | |
| def parse_args() -> argparse.Namespace: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| parser.add_argument("--config", type=Path, default=Path(__file__).resolve().parents[1] / "conf" / "config.yaml", help="Reserved for a consistent cluster interface; checkpoint config is authoritative") | |
| parser.add_argument("--checkpoint", type=Path, default=None, help="Checkpoint (default: ACE/data/checkpoint/model_bak.pt)") | |
| parser.add_argument("--input-path", type=Path, default=None, help="NPZ with inputs or initial_prognostic/forcings") | |
| parser.add_argument("--fake-data", action="store_true") | |
| parser.add_argument("--steps", type=int, default=4) | |
| parser.add_argument("--num-samples", type=int, default=1) | |
| parser.add_argument("--height", type=int, default=180) | |
| parser.add_argument("--width", type=int, default=360) | |
| parser.add_argument("--output-dir", type=Path, default=None, help="Inference output directory (default: ACE/output/infer)") | |
| parser.add_argument("--output-path", type=Path, default=None, help="Output NPZ path; overrides --output-dir/rollout.npz") | |
| parser.add_argument("--device", default="auto") | |
| return parser.parse_args() | |
| def main() -> int: | |
| args = parse_args() | |
| with args.config.open("r", encoding="utf-8") as handle: | |
| config = yaml.safe_load(handle) or {} | |
| checkpoint_path = args.checkpoint or configured_path(config, "checkpoint_path", CHECKPOINT_PATH) | |
| input_path = args.input_path or configured_path(config, "data_path", GENERATED_DATA_PATH) | |
| output_dir = args.output_dir or configured_path(config, "infer_dir", INFER_DIR) | |
| output_path = args.output_path or (output_dir / "rollout.npz") | |
| if not checkpoint_path.exists(): | |
| raise SystemExit( | |
| f"checkpoint not found: {checkpoint_path}; run 'python ACE/scripts/train.py' first" | |
| ) | |
| device_name = args.device | |
| if device_name == "auto": | |
| device_name = "cuda" if torch.cuda.is_available() else "cpu" | |
| checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False) | |
| model_values = dict(checkpoint["model_config"]) | |
| model_values.pop("modes_lat", None) | |
| model_values.pop("modes_lon", None) | |
| model_values["fallback"] = False | |
| model_config = ACEModelConfig(**model_values) | |
| model = ACEModel(model_config) | |
| state = checkpoint.get("ema_state", checkpoint["model_state"]) | |
| model.load_state_dict(state, strict=True) | |
| model.eval().to(device_name) | |
| normalizer = ACEDataNormalizer.from_dict(checkpoint["normalizer"]) if "normalizer" in checkpoint else None | |
| if args.fake_data: | |
| inputs, _ = make_fake_pairs(args.num_samples, args.height, args.width, seed=11) | |
| raw_inputs = inputs | |
| initial_np = raw_inputs[:, : model_config.prognostic_channels] | |
| forcing_np = np.repeat(raw_inputs[:, None, model_config.prognostic_channels :], args.steps, axis=1) | |
| source = "fake-data (smoke only)" | |
| else: | |
| if not input_path.exists(): | |
| data_cfg = config.get("data", {}) | |
| save_fake_pairs( | |
| input_path, | |
| num_samples=int(data_cfg.get("synthetic_num_samples", args.num_samples)), | |
| height=int(data_cfg.get("synthetic_height", args.height)), | |
| width=int(data_cfg.get("synthetic_width", args.width)), | |
| seed=0, | |
| ) | |
| data = np.load(input_path) | |
| if "initial_prognostic" in data and "forcings" in data: | |
| initial_np = np.asarray(data["initial_prognostic"], dtype=np.float32) | |
| forcing_np = np.asarray(data["forcings"], dtype=np.float32) | |
| elif "inputs" in data: | |
| initial_np = data["inputs"][:, : model_config.prognostic_channels] | |
| forcing_np = np.repeat(data["inputs"][:, None, model_config.prognostic_channels :], args.steps, axis=1) | |
| else: | |
| raise KeyError("input NPZ requires initial_prognostic/forcings or inputs") | |
| source = str(input_path) | |
| raw_initial_np = np.asarray(initial_np, dtype=np.float32) | |
| raw_forcing_np = np.asarray(forcing_np, dtype=np.float32) | |
| if normalizer is not None: | |
| repeated_state = np.repeat(raw_initial_np[:, None], raw_forcing_np.shape[1], axis=1) | |
| normalized = normalizer.transform_inputs( | |
| np.concatenate([repeated_state, raw_forcing_np], axis=2).reshape( | |
| -1, model_config.input_channels, raw_forcing_np.shape[-2], raw_forcing_np.shape[-1] | |
| ) | |
| ).reshape(repeated_state.shape[0], repeated_state.shape[1], model_config.input_channels, raw_forcing_np.shape[-2], raw_forcing_np.shape[-1]) | |
| initial_np = normalized[:, 0, : model_config.prognostic_channels] | |
| forcing_np = normalized[:, :, model_config.prognostic_channels :] | |
| initial = torch.from_numpy(initial_np).float().to(device_name) | |
| forcings = torch.from_numpy(forcing_np).float().to(device_name) | |
| with torch.no_grad(): | |
| output = model.rollout(initial, forcings, steps=args.steps).cpu().numpy() | |
| if normalizer is not None: | |
| output = normalizer.inverse_targets(output) | |
| output_path.parent.mkdir(parents=True, exist_ok=True) | |
| save_arrays = { | |
| "predictions": output, | |
| "initial_prognostic": raw_initial_np, | |
| "forcings": raw_forcing_np, | |
| } | |
| if not args.fake_data and "targets" in data: | |
| save_arrays["targets"] = np.asarray(data["targets"], dtype=np.float32) | |
| np.savez(output_path, **save_arrays) | |
| print(json.dumps({"status": "success", "output": str(output_path), "shape": list(output.shape), "source": source})) | |
| return 0 | |
| if __name__ == "__main__": | |
| raise SystemExit(main()) | |