Ai2_Climate_Emulator / scripts /inference.py
yzt15806542928's picture
Upload folder using huggingface_hub
5c365c5 verified
Raw
History Blame Contribute Delete
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())