File size: 10,297 Bytes
f4a39ee
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
#!/usr/bin/env python3
"""Run an official NeuralGCM pressure-level rollout."""
from __future__ import annotations
import argparse
import math
import pickle
import sys
import time
from pathlib import Path
import numpy as np
try:
    from common import PROJECT_ROOT, add_static_features, as_time_major_frames, era5_data_is_synthetic, era5_sample_to_xarray, load_config, load_era5_dataset, regrid_for_neuralgcm, resolve_path, validate_synthetic_era5_version
except ModuleNotFoundError:  # supports ``python -m scripts.inference``
    from scripts.common import PROJECT_ROOT, add_static_features, as_time_major_frames, era5_data_is_synthetic, era5_sample_to_xarray, load_config, load_era5_dataset, regrid_for_neuralgcm, resolve_path, validate_synthetic_era5_version
if str(PROJECT_ROOT) not in sys.path:
    sys.path.insert(0, str(PROJECT_ROOT))
from model.NeuralGCM import checkpoint_mode, load_checkpoint, validate_checkpoint_mode

MODE_ALIASES = {"forecast": "weather_forecast", "weather_forecast": "weather_forecast", "climate": "climate_scale", "climate_scale": "climate_scale", "forecast_2_8_deg": "forecast_2_8_deg", "stochastic_1_4_deg": "stochastic_1_4_deg"}


def _validate_checkpoint_mode(candidate: Path, mode: str) -> None:
    """Reject a checkpoint whose declared/profile grid differs from --mode."""
    try:
        with candidate.open("rb") as handle:
            payload = pickle.load(handle)
    except Exception:
        return
    validate_checkpoint_mode(payload, mode, candidate)


def _official_checkpoint(config: dict, mode: str, explicit: str | None) -> Path:
    if explicit:
        candidate = resolve_path(explicit)
        if candidate.exists():
            _validate_checkpoint_mode(candidate, mode)
        return candidate
    configured = config["inference"].get("checkpoint")
    if configured:
        candidate = resolve_path(configured)
        if candidate.exists():
            try:
                with candidate.open("rb") as handle:
                    payload = pickle.load(handle)
                if isinstance(payload, dict) and {"model_config_str", "aux_ds_dict", "params"}.issubset(payload):
                    try:
                        _validate_checkpoint_mode(candidate, mode)
                    except ValueError as exc:
                        # The default path is a convenience pointer. If a
                        # previous run left a checkpoint for another profile,
                        # use the bundled official checkpoint for the requested
                        # mode; explicit --checkpoint remains strict.
                        print(f"Warning: {exc}; falling back to profile checkpoint")
                        return resolve_path(config["model"]["profiles"][mode]["official_reference"])
                    return candidate
            except Exception:
                pass
        if Path(configured).name != "model_bak.pkl":
            return candidate
    return resolve_path(config["model"]["profiles"][mode]["official_reference"])


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--config", default="conf/config.yaml")
    parser.add_argument("--mode")
    parser.add_argument("--checkpoint", help="official NeuralGCM checkpoint (.pkl)")
    parser.add_argument("--input-nc")
    parser.add_argument("--data-dir")
    parser.add_argument("--years", nargs="+", type=int)
    parser.add_argument("--sample-index", type=int, default=0)
    parser.add_argument("--steps", type=int)
    parser.add_argument(
        "--output-interval-hours",
        type=int,
        help="hours between saved forecasts; defaults to inference.output_interval_hours",
    )
    parser.add_argument("--output")
    parser.add_argument(
        "--seed",
        type=int,
        help="PRNG seed for stochastic profiles; overrides inference.seed",
    )
    parser.add_argument("--device", default="auto", help="jax platform: auto, cpu, or gpu")
    args = parser.parse_args()
    config = load_config(args.config)
    if args.data_dir:
        config["data"]["data_dir"] = args.data_dir
        paired_static = resolve_path(args.data_dir, args.config) / "static.nc"
        if paired_static.exists():
            config["data"]["static_file"] = str(paired_static)
    requested = args.mode or config["inference"].get("mode", "weather_forecast")
    mode = MODE_ALIASES.get(requested, requested)
    if mode not in config["model"].get("profiles", {}):
        raise ValueError(f"Unknown mode {requested!r}")
    if args.device != "auto":
        import jax
        jax.config.update("jax_platform_name", args.device)
    checkpoint = _official_checkpoint(config, mode, args.checkpoint)
    model = load_checkpoint(checkpoint)
    print(f"Official checkpoint: {checkpoint}")
    grid = (model.data_coords.horizontal.latitudes.size, model.data_coords.horizontal.longitudes.size)
    print(f"Model mode={mode}, timestep={model.timestep}, grid={grid}")
    steps = int(args.steps or config["inference"].get("prediction_steps", 8))
    if steps <= 0:
        raise ValueError("steps must be positive")
    output_interval_hours = int(
        args.output_interval_hours
        or config["inference"].get("output_interval_hours", 6)
    )
    if output_interval_hours <= 0:
        raise ValueError("output-interval-hours must be positive")
    output_interval = np.timedelta64(output_interval_hours, "h")
    if args.input_nc:
        import xarray as xr
        dataset = xr.load_dataset(resolve_path(args.input_nc))
        synthetic_input = False
    else:
        years = args.years or list(config["data"].get("test_years", [2002]))
        synthetic_input = era5_data_is_synthetic(config, years)
        if synthetic_input:
            validate_synthetic_era5_version(config, years)
        data_interval_hours = int(config["data"].get("time_step_hours", 6))
        forcing_steps = math.ceil(steps * output_interval_hours / data_interval_hours)
        # Keep the forcing trajectory available at every requested lead time.
        # ERA5Dataset returns (initial state, future frames, ...); using its
        # default output_steps=1 would silently hold SST/sea-ice forcing fixed.
        source = load_era5_dataset(
            config, years, input_steps=1, output_steps=forcing_steps
        )
        # OneScience exposes ``total_samples`` directly, but its ``__len__``
        # returns the raw (possibly negative) window count.  Calling len() on
        # an undersized file therefore raises Python's own ``ValueError``
        # before we can explain which forcing window is missing.
        dataset_size = int(getattr(source, "total_samples", 0))
        if dataset_size <= 0:
            raise ValueError(
                "ERA5Dataset has no complete inference window: "
                f"T={getattr(source, 'T', '?')}, input_steps=1, "
                f"output_steps={forcing_steps}. Provide at least "
                f"{forcing_steps + 1} consecutive frames."
            )
        if not 0 <= args.sample_index < dataset_size:
            raise IndexError(
                f"sample-index {args.sample_index} outside dataset of length "
                f"{dataset_size}"
            )
        sample = source[args.sample_index]
        target_frames = as_time_major_frames(sample[1], name="ERA5 target")
        frames = [sample[0]] + [
            target_frames[t] for t in range(min(forcing_steps, len(target_frames)))
        ]
        import xarray as xr
        sample_time = sample[4][0]
        if not isinstance(sample_time, str) or len(sample_time) != 10 or not sample_time.isdigit():
            raise ValueError(f"invalid ERA5Dataset time index {sample_time!r}")
        start_timestamp = np.datetime64(
            f"{sample_time[:4]}-{sample_time[4:6]}-{sample_time[6:8]}T{sample_time[8:10]}:00:00"
        )
        dataset = xr.concat(
            [
                era5_sample_to_xarray(
                    (frame, frame),
                    config,
                    timestamp=start_timestamp
                    + np.timedelta64(int(config["data"].get("time_step_hours", 6)) * t, "h"),
                )
                for t, frame in enumerate(frames)
            ],
            dim="time",
        )
    dataset = regrid_for_neuralgcm(dataset, model)
    dataset = add_static_features(
        dataset, config, mode=mode, prefer_profile=not synthetic_input
    )
    data_in, forcings_in = model.data_from_xarray(dataset.isel(time=0))
    forcing_trajectory = model.forcings_from_xarray(dataset)
    import jax
    seed = int(
        args.seed
        if args.seed is not None
        else config["inference"].get("seed", config["project"].get("seed", 0))
    )
    state = model.encode(data_in, forcings_in, rng_key=jax.random.key(seed))
    start = time.perf_counter()
    _, outputs = model.unroll(
        state,
        forcing_trajectory,
        steps=steps,
        timedelta=output_interval,
    )
    times = np.arange(
        dataset.time.values[0] + output_interval,
        dataset.time.values[0] + (steps + 1) * output_interval,
        output_interval,
    )
    result = model.data_to_xarray(outputs, times=times)
    nonfinite = [
        name
        for name, values in result.data_vars.items()
        if np.issubdtype(values.dtype, np.number)
        and not bool(np.isfinite(values.values).all())
    ]
    if nonfinite:
        raise FloatingPointError(
            "NeuralGCM rollout produced NaN/Inf in variables "
            f"{nonfinite}. The input passed shape checks but is not a stable "
            "physical initial condition for this checkpoint."
        )
    result.attrs.update({"neuralgcm_mode": mode, "checkpoint": str(checkpoint), "official_api": "PressureLevelModel", "random_seed": seed})
    output = resolve_path(args.output or config["inference"].get("output", "results/predictions.nc"), args.config)
    output.parent.mkdir(parents=True, exist_ok=True)
    result.to_netcdf(output)
    print(
        f"Official rollout steps={steps}, output_interval={output_interval_hours}h, "
        f"horizon={steps * output_interval_hours / 24:g} days, "
        f"elapsed={time.perf_counter() - start:.2f}s"
    )
    print(f"Saved {output} with variables={list(result.data_vars)}")


if __name__ == "__main__":
    main()