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