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