File size: 6,200 Bytes
5c365c5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
"""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())