| """Official NeuralGCM implementation facade. |
| |
| The upstream legacy model, encoders, decoders, dynamical core and reference |
| training utilities are vendored directly under this project's ``model`` |
| namespace (``model/legacy`` and ``model/reference_code``). This file is the |
| single project-facing entry point; no external ``neuralgcm`` source directory |
| is required at runtime. |
| """ |
| from __future__ import annotations |
|
|
| import pickle |
| from pathlib import Path |
| from typing import Any |
|
|
| import numpy as np |
|
|
| PROFILE_GIN = { |
| "weather_forecast": "deterministic_0_7_deg.gin", |
| "climate_scale": "deterministic_1_4_deg.gin", |
| "forecast_2_8_deg": "deterministic_2_8_deg.gin", |
| "stochastic_1_4_deg": "stochastic_1_4_deg.gin", |
| } |
|
|
| 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", |
| } |
|
|
|
|
| class OfficialNeuralGCMUnavailable(RuntimeError): |
| """Raised when the official runtime package is not available.""" |
|
|
|
|
| class CheckpointFormatError(ValueError): |
| """Raised when a file is not an official NeuralGCM checkpoint.""" |
|
|
|
|
| def checkpoint_mode(payload: object) -> str | None: |
| """Infer the project profile declared by an official-format checkpoint.""" |
| if not isinstance(payload, dict): |
| return None |
| if payload.get("mode"): |
| value = str(payload["mode"]) |
| return MODE_ALIASES.get(value, value) |
| text = str(payload.get("model_config_str", "")) |
| if "GridTL255" in text: |
| return "weather_forecast" |
| if "GridTL63" in text: |
| return "forecast_2_8_deg" |
| if "GridTL127" in text: |
| return "stochastic_1_4_deg" if "FIELD_SUBSET" in text else "climate_scale" |
| return None |
|
|
|
|
| def validate_checkpoint_mode(payload: object, mode: str, path: str | Path) -> None: |
| """Reject a checkpoint whose grid/profile differs from the requested mode.""" |
| stored_mode = checkpoint_mode(payload) |
| if stored_mode and stored_mode != mode: |
| raise ValueError( |
| f"Checkpoint {path} is for mode={stored_mode!r}, but mode={mode!r} " |
| "was requested. Select the matching mode or checkpoint." |
| ) |
|
|
|
|
| def parameter_summary(params: Any) -> dict[str, Any]: |
| """Return reproducible parameter count, storage size and dtype statistics.""" |
| import jax |
|
|
| leaves = jax.tree_util.tree_leaves(params) |
| array_leaves = [leaf for leaf in leaves if hasattr(leaf, "shape") and hasattr(leaf, "dtype")] |
| count = sum(int(np.prod(leaf.shape, dtype=np.int64)) for leaf in array_leaves) |
| nbytes = sum( |
| int(np.prod(leaf.shape, dtype=np.int64)) * np.dtype(leaf.dtype).itemsize |
| for leaf in array_leaves |
| ) |
| dtype_counts: dict[str, int] = {} |
| for leaf in array_leaves: |
| dtype = str(np.dtype(leaf.dtype)) |
| dtype_counts[dtype] = dtype_counts.get(dtype, 0) + int( |
| np.prod(leaf.shape, dtype=np.int64) |
| ) |
| return { |
| "count": count, |
| "nbytes": nbytes, |
| "leaves": len(array_leaves), |
| "dtypes": dtype_counts, |
| } |
|
|
|
|
| def format_parameter_summary(params: Any) -> str: |
| """Format a compact ``params.count``-style model summary.""" |
| summary = parameter_summary(params) |
| dtype_text = ",".join( |
| f"{dtype}:{count:,}" for dtype, count in sorted(summary["dtypes"].items()) |
| ) |
| return ( |
| f"params.count={summary['count']:,} " |
| f"params.bytes={summary['nbytes']:,} " |
| f"params.mib={summary['nbytes'] / 2**20:.2f} " |
| f"params.leaves={summary['leaves']} dtypes={dtype_text}" |
| ) |
|
|
|
|
| def load_checkpoint(path: str | Path): |
| """Load an official checkpoint through the vendored PressureLevelModel.""" |
| try: |
| from model.legacy.api import PressureLevelModel |
| except Exception as exc: |
| raise OfficialNeuralGCMUnavailable( |
| "Unable to import the vendored NeuralGCM implementation. Check " |
| "JAX, Haiku, Gin and Dinosaur dependencies in develop_base." |
| ) from exc |
| path = Path(path) |
| if not path.exists(): |
| raise FileNotFoundError(path) |
| with path.open("rb") as handle: |
| checkpoint = pickle.load(handle) |
| required = {"model_config_str", "aux_ds_dict", "params"} |
| if not isinstance(checkpoint, dict) or not required.issubset(checkpoint): |
| keys = sorted(checkpoint) if isinstance(checkpoint, dict) else type(checkpoint).__name__ |
| raise CheckpointFormatError( |
| f"{path} is not an official checkpoint; expected keys " |
| f"{sorted(required)}, got {keys}" |
| ) |
| return PressureLevelModel.from_checkpoint(checkpoint) |
|
|
|
|
| def official_runtime_available() -> bool: |
| try: |
| from model.legacy.api import PressureLevelModel |
| except Exception: |
| return False |
| return True |
|
|
|
|
| def build_from_scratch(dataset, mode: str): |
| """Build the public WhirlModel used for random parameter initialization. |
| |
| Parameter initialization itself needs a concrete trajectory and is performed |
| by ``scripts/train.py`` through the returned model's Haiku rollout function. |
| This compatibility facade deliberately does not import the unreleased Google |
| experiment runner. |
| """ |
| return build_training_model(dataset, mode) |
|
|
|
|
| def build_training_model(dataset, mode: str): |
| """Build an official ``WhirlModel`` from the fused Gin profile.""" |
| if mode not in PROFILE_GIN: |
| raise ValueError(f"Unknown NeuralGCM mode {mode!r}") |
| import gin |
| from model.legacy import model_builder |
|
|
| config_path = Path(__file__).resolve().parent / "reference_code" / "paper_configs" / PROFILE_GIN[mode] |
| gin_text = config_path.read_text(encoding="utf-8") |
| |
| |
| |
| |
| from dinosaur import xarray_utils |
| try: |
| aux_features = xarray_utils.aux_features_from_xarray(dataset) |
| except (KeyError, AttributeError): |
| aux_features = {} |
| aux_features[xarray_utils.XARRAY_DS_KEY] = dataset |
| dataset = dataset.copy() |
| dataset.attrs = dict(dataset.attrs) |
| dataset.attrs[xarray_utils.XR_AUX_FEATURES_LIST_KEY] = ",".join( |
| key for key in aux_features if key != xarray_utils.XARRAY_DS_KEY |
| ) |
| |
| |
| original = model_builder.xarray_utils.aux_features_from_xarray |
| model_builder.xarray_utils.aux_features_from_xarray = lambda _: aux_features |
| try: |
| model = model_builder.get_whirl_model(dataset, gin_text) |
| finally: |
| model_builder.xarray_utils.aux_features_from_xarray = original |
| |
| |
| return model, gin_text |
|
|
|
|
| def make_rollout_functions( |
| whirl_model, trajectory_length: int, *, inner_steps: int = 1 |
| ): |
| """Return Haiku init/apply functions using the official rollout helpers.""" |
| import haiku as hk |
| from model.legacy import model_utils |
|
|
| @hk.transform |
| def rollout_fn(target, forcing): |
| model = whirl_model.model_cls() |
| trajectory_fn = model_utils.trajectory_with_inputs_and_forcing( |
| model, num_init_frames=1, start_with_input=True |
| ) |
| _, predicted = trajectory_fn( |
| target, |
| forcing, |
| outer_steps=trajectory_length, |
| inner_steps=inner_steps, |
| ) |
| return model_utils.compute_prediction_and_target_representations( |
| predicted, target, forcing, model |
| ) |
|
|
| return rollout_fn |
|
|
|
|
| def save_official_checkpoint(path: str | Path, params: Any, dataset, model_config_str: str, *, metadata: dict[str, Any] | None = None): |
| """Write a checkpoint consumable by ``PressureLevelModel.from_checkpoint``.""" |
| path = Path(path) |
| path.parent.mkdir(parents=True, exist_ok=True) |
| payload = { |
| "model_config_str": model_config_str, |
| "aux_ds_dict": dataset.to_dict(), |
| "params": params, |
| } |
| if metadata: |
| payload.update(metadata) |
| with path.open("wb") as handle: |
| pickle.dump(payload, handle, protocol=pickle.HIGHEST_PROTOCOL) |
| return path |
|
|
|
|
| NeuralGCMAdapter = load_checkpoint |
|
|