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