| |
| """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: |
| 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: |
| |
| |
| |
| |
| 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) |
| |
| |
| |
| source = load_era5_dataset( |
| config, years, input_steps=1, output_steps=forcing_steps |
| ) |
| |
| |
| |
| |
| 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() |
|
|