| |
| """Train official NeuralGCM dynamics with a OneScience ERA5Dataset source.""" |
| from __future__ import annotations |
|
|
| import argparse |
| from collections.abc import Mapping |
| import pickle |
| import sys |
| import time |
| from pathlib import Path |
|
|
| import numpy as np |
| import optax |
|
|
| try: |
| from common import PROJECT_ROOT, add_static_features, as_time_major_frames, era5_data_is_synthetic, era5_frames_to_xarray, load_config, load_era5_dataset, regrid_for_profile, 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_frames_to_xarray, load_config, load_era5_dataset, regrid_for_profile, 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 build_training_model, format_parameter_summary, make_rollout_functions, parameter_summary, save_official_checkpoint, validate_checkpoint_mode |
| try: |
| from losses import make_loss_fn |
| except ModuleNotFoundError: |
| from scripts.losses import make_loss_fn |
| try: |
| from streaming_data import PrefetchedWindowBatches, WindowBatchStream |
| except ModuleNotFoundError: |
| from scripts.streaming_data import PrefetchedWindowBatches, WindowBatchStream |
|
|
| 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 _merge_training_profile( |
| training: Mapping, mode: str, *, paper_defaults: bool = False |
| ) -> dict: |
| """Merge mode semantics and optionally configured long-run settings.""" |
| merged = dict(training) |
| profiles = merged.pop("profiles", {}) |
| profile = dict(profiles.get(mode, {})) |
| if not paper_defaults: |
| profile = { |
| key: value |
| for key, value in profile.items() |
| if key in {"ensemble_size", "loss"} |
| } |
| for key, value in profile.items(): |
| if isinstance(value, Mapping) and isinstance(merged.get(key), Mapping): |
| merged[key] = {**merged[key], **value} |
| else: |
| merged[key] = value |
| return merged |
|
|
|
|
| def _make_learning_rate_schedule(optimizer_cfg: Mapping, peak_rate: float): |
| """Build a configured NeuralGCM-style or constant schedule.""" |
| schedule_name = str(optimizer_cfg.get("schedule", "constant")).lower() |
| if schedule_name == "neuralgcm": |
| warmup_steps = int(optimizer_cfg.get("warmup_steps", 2000)) |
| decay_start = int(optimizer_cfg.get("decay_start", 15000)) |
| decay_steps = int(optimizer_cfg.get("decay_steps", 10000)) |
| decay_rate = float(optimizer_cfg.get("decay_rate", 0.5)) |
| if warmup_steps <= 0 or decay_steps <= 0 or decay_start < warmup_steps: |
| raise ValueError("invalid NeuralGCM optimizer schedule boundaries") |
| warmup = optax.linear_schedule(0.0, peak_rate, warmup_steps) |
| plateau = optax.constant_schedule(peak_rate) |
| decay = optax.exponential_decay( |
| peak_rate, |
| transition_steps=decay_steps, |
| decay_rate=decay_rate, |
| staircase=False, |
| ) |
| return optax.join_schedules( |
| (warmup, plateau, decay), (warmup_steps, decay_start) |
| ) |
| if schedule_name != "constant": |
| raise ValueError(f"Unknown training.optimizer.schedule {schedule_name!r}") |
| rates = [float(x) for x in optimizer_cfg.get("rates", [])] |
| boundaries = [int(x) for x in optimizer_cfg.get("boundaries", [])] |
| if rates: |
| if len(rates) != len(boundaries) + 1: |
| raise ValueError("training.optimizer.rates must have one more entry than boundaries") |
| return optax.join_schedules( |
| [optax.constant_schedule(rate) for rate in rates], boundaries |
| ) |
| return optax.constant_schedule(peak_rate) |
|
|
|
|
| def _era5_frame_capacity(config: Mapping, years: list[int]) -> tuple[int | None, bool]: |
| """Return the shortest yearly trajectory and whether all files are virtual.""" |
| import h5py |
|
|
| data_dir = resolve_path(config["data"]["data_dir"]) / "data" |
| paths = [data_dir / f"{year}.h5" for year in years] |
| if not paths or any(not path.exists() for path in paths): |
| return None, False |
| frame_counts = [] |
| synthetic_flags = [] |
| try: |
| for path in paths: |
| with h5py.File(path, "r") as handle: |
| fields = handle[config["data"].get("field_key", "fields")] |
| frame_counts.append(int(fields.shape[0])) |
| synthetic_flags.append(bool(fields.attrs.get("synthetic", False))) |
| except (KeyError, OSError): |
| |
| return None, False |
| return min(frame_counts), all(synthetic_flags) |
|
|
|
|
| def _fit_rollout_schedule_to_data( |
| rollout_schedule: list[dict], |
| *, |
| available_frames: int | None, |
| synthetic: bool, |
| explicit_override: bool, |
| ) -> list[dict]: |
| """Fit long rollouts to virtual data without weakening real-data checks.""" |
| required_frames = max(item["trajectory_length"] for item in rollout_schedule) |
| if available_frames is None or required_frames <= available_frames: |
| return rollout_schedule |
| if explicit_override or not synthetic: |
| source = "explicit --trajectory-length" if explicit_override else "training profile" |
| raise ValueError( |
| f"{source} requires {required_frames} consecutive ERA5 frames, but " |
| f"the shortest training-year file has {available_frames}. Generate a " |
| "longer trajectory or lower --trajectory-length." |
| ) |
| fitted = [ |
| item for item in rollout_schedule |
| if item["trajectory_length"] <= available_frames |
| ] |
| if not fitted and available_frames >= 2: |
| fitted = [{"trajectory_length": available_frames, "until_step": 0}] |
| if not fitted: |
| raise ValueError( |
| "Virtual ERA5 data needs at least two consecutive frames for training; " |
| f"found {available_frames}." |
| ) |
| print( |
| f"data virtual_rollout_clamped={required_frames}->{max(item['trajectory_length'] for item in fitted)} frames" |
| ) |
| return fitted |
|
|
|
|
| def _trajectory_from_dataset(dataset, steps: int): |
| """Convert a time-indexed xarray sample to official model dictionaries.""" |
| import gin |
| from model.legacy import model_builder |
| |
| state_fn = gin.query_parameter("WhirlModel.from_xarray_fn") |
| del state_fn |
| converter = model_builder.xarray_to_state_and_dynamic_covariate_data |
| state_data, forcing_data = converter(dataset) |
| return state_data, forcing_data |
|
|
|
|
| def _replicate_tree(tree, devices): |
| """Replicate a pytree along a leading local-device axis.""" |
| import jax |
|
|
| return jax.tree_util.tree_map( |
| lambda value: jax.device_put_replicated(value, devices), tree |
| ) |
|
|
|
|
| def _unreplicate_tree(tree): |
| """Take replica zero back to host for a normal checkpoint.""" |
| import jax |
|
|
| return jax.tree_util.tree_map(lambda value: jax.device_get(value[0]), tree) |
|
|
|
|
| def _stack_trees(trees): |
| """Stacks trajectory pytrees as host arrays, ready for direct sharding.""" |
| if not trees: |
| raise ValueError("cannot stack an empty pytree sequence") |
| import jax |
| return jax.tree_util.tree_map( |
| lambda *values: np.stack([np.asarray(value) for value in values], axis=0), |
| *trees, |
| ) |
|
|
|
|
| def _put_batch_sharded(tree, devices, local_batch): |
| """Place each host batch slice directly on its destination device.""" |
| import jax |
|
|
| device_count = len(devices) |
|
|
| def put(value): |
| value = np.asarray(value) |
| expected = device_count * local_batch |
| if value.shape[0] != expected: |
| raise ValueError( |
| f"batch leaf has leading size {value.shape[0]}, expected {expected}" |
| ) |
| value = value.reshape((device_count, local_batch) + value.shape[1:]) |
| return jax.device_put_sharded( |
| [value[index] for index in range(device_count)], devices |
| ) |
|
|
| return jax.tree_util.tree_map(put, tree) |
|
|
|
|
| def _sample_start_time(sample): |
| """Return the real first timestamp supplied by OneScience ERA5Dataset.""" |
| time_index = sample[4] |
| if not time_index: |
| raise ValueError("ERA5Dataset sample has an empty time_index") |
| value = str(time_index[0]) |
| if len(value) != 10 or not value.isdigit(): |
| raise ValueError(f"invalid ERA5Dataset time index {value!r}") |
| return np.datetime64( |
| f"{value[:4]}-{value[4:6]}-{value[6:8]}T{value[8:10]}:00:00" |
| ) |
|
|
|
|
| def _read_checkpoint(path_value: str, mode: str) -> tuple[Path, dict]: |
| path = resolve_path(path_value) |
| with path.open("rb") as handle: |
| payload = pickle.load(handle) |
| if not isinstance(payload, dict) or "params" not in payload: |
| raise ValueError(f"Checkpoint {path} does not contain NeuralGCM params") |
| validate_checkpoint_mode(payload, mode, path) |
| return path, payload |
|
|
|
|
| def _validate_resume_contract(saved: Mapping, current: Mapping) -> None: |
| """Reject changes that would invalidate restored optimizer/data state.""" |
| mismatches = { |
| key: (saved.get(key), value) |
| for key, value in current.items() |
| if saved.get(key) != value |
| } |
| if mismatches: |
| details = ", ".join( |
| f"{key}: saved={old!r}, current={new!r}" |
| for key, (old, new) in mismatches.items() |
| ) |
| raise ValueError(f"Resume checkpoint is incompatible with this run: {details}") |
|
|
|
|
| def train( |
| config: dict, |
| mode: str, |
| finetune: str | None, |
| max_steps: int | None, |
| learning_rate: float | None, |
| devices_requested: int | None = None, |
| data_workers_requested: int | None = None, |
| prefetch_batches_requested: int | None = None, |
| trajectory_length_requested: int | None = None, |
| checkpoint_output: str | None = None, |
| paper_defaults: bool = False, |
| resume: str | None = None, |
| checkpoint_interval_requested: int | None = None, |
| loss_backend_requested: str | None = None, |
| ): |
| import jax |
| import jax.numpy as jnp |
| train_cfg = _merge_training_profile( |
| config.get("training", {}), mode, paper_defaults=paper_defaults |
| ) |
| available_devices = jax.local_devices() |
| requested_devices = int( |
| devices_requested |
| if devices_requested is not None |
| else train_cfg.get("devices", 1) |
| ) |
| if requested_devices <= 0: |
| raise ValueError("--devices must be a positive integer") |
| if requested_devices > len(available_devices): |
| raise RuntimeError( |
| f"Requested {requested_devices} local devices, but JAX exposes " |
| f"only {len(available_devices)}: {available_devices}" |
| ) |
| devices = available_devices[:requested_devices] |
| configured_batch = max(1, int(train_cfg.get("samples_per_step", 1))) |
| |
| |
| |
| global_batch = max(configured_batch, requested_devices) |
| global_batch = ((global_batch + requested_devices - 1) // requested_devices) * requested_devices |
| years = list(config["data"].get("train_years", [2000])) |
| trajectory_length = max(1, int(train_cfg.get("trajectory_length", 2))) |
| rollout_schedule_cfg = ( |
| [] if trajectory_length_requested is not None else train_cfg.get("rollout_schedule", []) |
| ) |
| if trajectory_length_requested is not None: |
| trajectory_length = int(trajectory_length_requested) |
| if rollout_schedule_cfg: |
| rollout_schedule = sorted( |
| [ |
| { |
| "trajectory_length": max(2, int(item["trajectory_length"])), |
| "until_step": int(item.get("until_step", 0)), |
| } |
| for item in rollout_schedule_cfg |
| ], |
| key=lambda item: item["until_step"], |
| ) |
| if rollout_schedule[0]["until_step"] not in (0, 1): |
| raise ValueError("training.rollout_schedule must start at until_step 0 or 1") |
| else: |
| rollout_schedule = [{"trajectory_length": trajectory_length, "until_step": 0}] |
| available_frames, synthetic_data = _era5_frame_capacity(config, years) |
| if synthetic_data: |
| validate_synthetic_era5_version(config, years) |
| rollout_schedule = _fit_rollout_schedule_to_data( |
| rollout_schedule, |
| available_frames=available_frames, |
| synthetic=synthetic_data, |
| explicit_override=trajectory_length_requested is not None, |
| ) |
| trajectory_length = max(item["trajectory_length"] for item in rollout_schedule) |
| if trajectory_length < 2: |
| raise ValueError( |
| "training.trajectory_length must be at least 2 (one initial and " |
| "one future ERA5 frame)" |
| ) |
| |
| |
| |
| future_steps = trajectory_length - 1 |
| dataset = load_era5_dataset( |
| config, years, input_steps=1, output_steps=future_steps |
| ) |
| dataset_size = int(getattr(dataset, "total_samples", -1)) |
| if dataset_size < 0: |
| raise ValueError( |
| "OneScience ERA5Dataset computed a negative sample count: " |
| f"T={dataset.T}, input_steps={dataset.input_steps}, " |
| f"output_steps={dataset.output_steps}. The requested trajectory is " |
| "longer than the data file." |
| ) |
| print(f"data samples={dataset_size} shape={(dataset.C, dataset.H, dataset.W)}") |
| if dataset_size < global_batch: |
| raise ValueError( |
| f"Training requires global batch={global_batch} trajectories for " |
| f"{requested_devices} devices, but ERA5Dataset has only {dataset_size} " |
| "samples. Generate more windows or lower training.samples_per_step." |
| ) |
| first_sample = dataset[0] |
| first_input = first_sample[0] |
| first_targets = as_time_major_frames(first_sample[1], name="ERA5 target") |
| first_frames = np.concatenate((first_input[None, ...], first_targets), axis=0) |
| |
| |
| ds = regrid_for_profile( |
| era5_frames_to_xarray( |
| first_frames, config, start_time=_sample_start_time(first_sample) |
| ), |
| mode, |
| ) |
| ds = add_static_features( |
| ds, config, mode=mode, prefer_profile=not synthetic_data |
| ) |
| model, gin_text = build_training_model(ds, mode) |
| |
| |
| if model.from_xarray_fn is None: |
| raise RuntimeError("Gin profile did not configure WhirlModel.from_xarray_fn") |
| |
| |
| from dinosaur import xarray_utils |
| reference_datetime = model.specs.aux_features["reference_datetime"] |
| def convert_samples(samples): |
| """Convert already-prefetched OneScience samples on the main thread.""" |
| converted = [] |
| for sample in samples: |
| input_frame = sample[0] |
| target_frames = as_time_major_frames(sample[1], name="ERA5 target") |
| frame_arrays = np.concatenate( |
| (np.asarray(input_frame)[None, ...], target_frames), axis=0 |
| ) |
| sample_ds = regrid_for_profile( |
| era5_frames_to_xarray( |
| frame_arrays, |
| config, |
| start_time=_sample_start_time(sample), |
| ), |
| mode, |
| ) |
| sample_ds = add_static_features( |
| sample_ds, |
| config, |
| mode=mode, |
| prefer_profile=not synthetic_data, |
| ) |
| sample_ds = xarray_utils.ds_with_sim_time( |
| sample_ds, |
| model.specs.physics_specs, |
| reference_datetime=reference_datetime, |
| ) |
| converted.append(model.from_xarray_fn(sample_ds)) |
| return ( |
| _stack_trees([item[0] for item in converted]), |
| _stack_trees([item[1] for item in converted]), |
| ) |
|
|
| |
| |
| initial_target, initial_forcing = convert_samples([first_sample]) |
| target = jax.tree_util.tree_map(lambda value: value[0], initial_target) |
| forcing_data = jax.tree_util.tree_map(lambda value: value[0], initial_forcing) |
| |
| |
| |
| data_interval = np.timedelta64( |
| int(config["data"].get("time_step_hours", 6)), "h" |
| ) |
| model_timestep = model.specs.physics_specs.dimensionalize_timedelta64( |
| model.specs.dt |
| ) |
| ratio = data_interval / model_timestep |
| inner_steps = int(round(float(ratio))) |
| if inner_steps <= 0 or abs(float(ratio) - inner_steps) > 1e-6: |
| raise ValueError( |
| f"ERA5 interval {data_interval} is not an integer multiple of " |
| f"NeuralGCM timestep {model_timestep}" |
| ) |
| rollout_max = make_rollout_functions( |
| model, trajectory_length=trajectory_length, inner_steps=inner_steps |
| ) |
| if finetune and resume: |
| raise ValueError("--finetune and --resume are mutually exclusive") |
| resume_path = None |
| resume_state = None |
| params = None |
| if resume: |
| resume_path, resume_payload = _read_checkpoint(resume, mode) |
| resume_state = resume_payload.get("training_state") |
| if not isinstance(resume_state, Mapping): |
| raise ValueError( |
| f"--resume requires a project checkpoint with full training_state; " |
| f"{resume_path} is inference-only. Use --finetune to load params only." |
| ) |
| if int(resume_state.get("format_version", -1)) != 1: |
| raise ValueError( |
| f"Unsupported training_state format in {resume_path}: " |
| f"{resume_state.get('format_version')!r}" |
| ) |
| params = resume_state.get("train_params") |
| if params is None: |
| raise ValueError(f"Resume checkpoint {resume_path} has no train_params") |
| if finetune: |
| _, payload = _read_checkpoint(finetune, mode) |
| params = payload["params"] |
| if params is None: |
| params = rollout_max.init(jax.random.key(int(config["project"].get("seed", 0))), target, forcing_data) |
| print(f"model mode={mode} {format_parameter_summary(params)}") |
| effective_lr = float( |
| learning_rate if learning_rate is not None else train_cfg.get("learning_rate", 1e-4) |
| ) |
| clip_norm = float(train_cfg.get("gradient_clip_norm", 1.0)) |
| optimizer_cfg = dict(train_cfg.get("optimizer", {})) |
| schedule = _make_learning_rate_schedule(optimizer_cfg, effective_lr) |
| b1 = float(optimizer_cfg.get("b1", 0.9)) |
| b2 = float(optimizer_cfg.get("b2", 0.95)) |
| eps = float(optimizer_cfg.get("eps", 1e-6)) |
| if clip_norm > 0: |
| optimizer = optax.chain( |
| optax.clip_by_global_norm(clip_norm), |
| optax.adam(schedule, b1=b1, b2=b2, eps=eps), |
| ) |
| else: |
| optimizer = optax.adam(schedule, b1=b1, b2=b2, eps=eps) |
| opt_state = optimizer.init(params) |
| if resume_state is not None: |
| restored_opt_state = resume_state.get("opt_state") |
| if restored_opt_state is None: |
| raise ValueError(f"Resume checkpoint {resume_path} has no opt_state") |
| opt_state = restored_opt_state |
|
|
| loss_config = dict(train_cfg.get("loss", {})) |
| if loss_backend_requested is not None: |
| loss_config["backend"] = loss_backend_requested |
| loss_backend = str(loss_config.get("backend", "official")).lower() |
| crps_training = loss_backend == "crps" |
| ensemble_size = int(train_cfg.get("ensemble_size", 2 if crps_training else 1)) |
| if crps_training and ensemble_size != 2: |
| raise ValueError("Official NeuralGCM CRPS training requires ensemble_size=2") |
| if not crps_training and ensemble_size != 1: |
| raise ValueError("Deterministic training requires ensemble_size=1") |
| rollout_cache = {trajectory_length: rollout_max} |
| loss_cache = { |
| trajectory_length: make_loss_fn( |
| model, |
| steps_per_save=inner_steps, |
| trajectory_length=trajectory_length, |
| config=loss_config, |
| mode=mode, |
| ) |
| } |
|
|
| def schedule_length(step: int) -> int: |
| selected = rollout_schedule[0]["trajectory_length"] |
| |
| for item in rollout_schedule[1:]: |
| if step > item["until_step"]: |
| selected = item["trajectory_length"] |
| return selected |
|
|
| def _slice_time(tree, length: int): |
| """Slice only trajectory leaves while retaining static metadata.""" |
| def slice_leaf(value): |
| shape = getattr(value, "shape", ()) |
| if len(shape) >= 2 and shape[1] == trajectory_length: |
| return value[:, :length] |
| if len(shape) and shape[0] == trajectory_length: |
| return value[:length] |
| return value |
| return jax.tree_util.tree_map(slice_leaf, tree) |
|
|
| def get_rollout_and_loss(length: int): |
| if length not in rollout_cache: |
| rollout_cache[length] = make_rollout_functions( |
| model, trajectory_length=length, inner_steps=inner_steps |
| ) |
| loss_cache[length] = make_loss_fn( |
| model, |
| steps_per_save=inner_steps, |
| trajectory_length=length, |
| config=loss_config, |
| mode=mode, |
| ) |
| return rollout_cache[length], loss_cache[length] |
| ema_num_steps = int(train_cfg.get("ema_num_steps", 0)) |
| ema_decay = 0.0 if ema_num_steps <= 0 else 1.0 - 2.0 / (ema_num_steps + 1.0) |
| resume_contract = { |
| "mode": mode, |
| "dataset_size": dataset_size, |
| "global_batch": global_batch, |
| "trajectory_length": trajectory_length, |
| "rollout_schedule": rollout_schedule, |
| "inner_steps": inner_steps, |
| "data_interval_hours": int(config["data"].get("time_step_hours", 6)), |
| "learning_rate": effective_lr, |
| "gradient_clip_norm": clip_norm, |
| "optimizer": optimizer_cfg, |
| "ema_num_steps": ema_num_steps, |
| "loss": loss_config, |
| "ensemble_size": ensemble_size, |
| "paper_defaults": paper_defaults, |
| } |
| start_step = 0 |
| restored_ema_params = None |
| if resume_state is not None: |
| saved_contract = resume_state.get("contract") |
| if not isinstance(saved_contract, Mapping): |
| raise ValueError(f"Resume checkpoint {resume_path} has no contract") |
| _validate_resume_contract(saved_contract, resume_contract) |
| start_step = int(resume_state.get("step", -1)) |
| if start_step < 0: |
| raise ValueError(f"Resume checkpoint {resume_path} has invalid step={start_step}") |
| restored_ema_params = resume_state.get("ema_params") |
| if restored_ema_params is None: |
| raise ValueError(f"Resume checkpoint {resume_path} has no ema_params") |
| stream = WindowBatchStream( |
| size=dataset_size, |
| global_batch=global_batch, |
| seed=int(config["project"].get("seed", 0)), |
| shuffle=bool(train_cfg.get("shuffle", True)), |
| drop_last=bool(train_cfg.get("drop_last", True)), |
| ) |
| if resume_state is not None: |
| data_stream_state = resume_state.get("data_stream_state") |
| if not isinstance(data_stream_state, dict): |
| raise ValueError( |
| f"Resume checkpoint {resume_path} has no data_stream_state" |
| ) |
| stream.load_state_dict(data_stream_state) |
| prefetcher = PrefetchedWindowBatches( |
| dataset, |
| stream, |
| num_workers=int( |
| data_workers_requested |
| if data_workers_requested is not None |
| else train_cfg.get("data_num_workers", 2) |
| ), |
| prefetch_batches=int( |
| prefetch_batches_requested |
| if prefetch_batches_requested is not None |
| else train_cfg.get("prefetch_batches", 1) |
| ), |
| ) |
|
|
| def loss_fn_for_length( |
| length, p, rngs, xs, fs, *, device_axis_name=None |
| ): |
| rollout_fn, trajectory_loss = get_rollout_and_loss(length) |
| def single_rollout(rng, x, f): |
| pred, truth = rollout_fn.apply(p, rng, x, f) |
| return pred, truth |
|
|
| if crps_training: |
| def single_ensemble_loss(member_rngs, x, f): |
| predictions, targets = jax.vmap( |
| single_rollout, |
| in_axes=(0, None, None), |
| spmd_axis_name="ensemble", |
| )(member_rngs, x, f) |
| per_member = jax.vmap( |
| trajectory_loss, |
| axis_name="ensemble", |
| spmd_axis_name="ensemble", |
| )(predictions, targets) |
| return jnp.mean(per_member) |
|
|
| per_example = jax.vmap( |
| single_ensemble_loss, |
| in_axes=(0, 0, 0), |
| axis_name="batch", |
| spmd_axis_name="batch", |
| )(rngs, xs, fs) |
| else: |
| predictions, targets = jax.vmap( |
| single_rollout, in_axes=(0, 0, 0) |
| )(rngs, xs, fs) |
| if hasattr(trajectory_loss, "evaluate_batch"): |
| return trajectory_loss.evaluate_batch( |
| predictions, |
| targets, |
| device_axis_name=device_axis_name, |
| ) |
| per_example = jax.vmap(trajectory_loss, in_axes=(0, 0))( |
| predictions, targets |
| ) |
| return jnp.mean(per_example) |
|
|
| def step_rngs(step: int, batch_size: int): |
| base_key = jax.random.key(int(config["project"].get("seed", 0))) |
| keys = jax.random.split( |
| jax.random.fold_in(base_key, step), batch_size * ensemble_size |
| ) |
| if crps_training: |
| return keys.reshape((batch_size, ensemble_size) + keys.shape[1:]) |
| return keys |
|
|
| if requested_devices == 1: |
| train_params, train_opt_state = params, opt_state |
| train_ema_params = ( |
| restored_ema_params |
| if restored_ema_params is not None |
| else jax.tree_util.tree_map(lambda value: value, params) |
| ) |
| value_grad_cache = {} |
|
|
| def train_step(step, current_params, current_opt_state, current_ema_params, batch_target, batch_forcing): |
| length = schedule_length(step) |
| if length not in value_grad_cache: |
| value_grad_cache[length] = jax.jit(jax.value_and_grad( |
| lambda p, r, x, f: loss_fn_for_length( |
| length, p, r, x, f |
| ) |
| )) |
| value_grad = value_grad_cache[length] |
| batch_target = _slice_time(batch_target, length) |
| batch_forcing = _slice_time(batch_forcing, length) |
| rngs = step_rngs(step, global_batch) |
| loss, grads = value_grad( |
| current_params, rngs, batch_target, batch_forcing |
| ) |
| grad_finite = np.asarray( |
| jax.device_get( |
| jnp.asarray( |
| [ |
| jnp.all(jnp.isfinite(g)) |
| for g in jax.tree_util.tree_leaves(grads) |
| ] |
| ) |
| ) |
| ) |
| loss_value = float(np.asarray(jax.device_get(loss))) |
| if not np.isfinite(loss_value) or not np.all(grad_finite): |
| bad_grad_leaves = int(np.size(grad_finite) - np.count_nonzero(grad_finite)) |
| raise FloatingPointError( |
| f"NeuralGCM produced loss={loss_value!r} and " |
| f"nonfinite_gradient_leaves={bad_grad_leaves}; " |
| "reduce learning_rate, increase gradient clipping, or " |
| "use a physically consistent ERA5 trajectory." |
| ) |
| updates, current_opt_state = optimizer.update( |
| grads, current_opt_state, current_params |
| ) |
| current_params = optax.apply_updates(current_params, updates) |
| current_ema_params = jax.tree_util.tree_map( |
| lambda old, new: ema_decay * old + (1.0 - ema_decay) * new, |
| current_ema_params, |
| current_params, |
| ) |
| return current_params, current_opt_state, current_ema_params, loss |
|
|
| else: |
| |
| |
| local_batch = global_batch // requested_devices |
| train_params = _replicate_tree(params, devices) |
| train_opt_state = _replicate_tree(opt_state, devices) |
| train_ema_params = _replicate_tree( |
| restored_ema_params if restored_ema_params is not None else params, |
| devices, |
| ) |
|
|
| def make_pmapped_step(length): |
| def pmapped_step(current_params, current_opt_state, current_ema_params, rng, x, f): |
| loss, grads = jax.value_and_grad( |
| lambda p, r, xx, ff: loss_fn_for_length( |
| length, |
| p, |
| r, |
| xx, |
| ff, |
| device_axis_name="devices", |
| ) |
| )(current_params, rng, x, f) |
| grads = jax.lax.pmean(grads, axis_name="devices") |
| loss = jax.lax.pmean(loss, axis_name="devices") |
| grad_finite = jnp.all( |
| jnp.asarray( |
| [jnp.all(jnp.isfinite(g)) for g in jax.tree_util.tree_leaves(grads)] |
| ) |
| ) |
| updates, current_opt_state = optimizer.update( |
| grads, current_opt_state, current_params |
| ) |
| current_params = optax.apply_updates(current_params, updates) |
| current_ema_params = jax.tree_util.tree_map( |
| lambda old, new: ema_decay * old + (1.0 - ema_decay) * new, |
| current_ema_params, |
| current_params, |
| ) |
| return current_params, current_opt_state, current_ema_params, loss, grad_finite |
| return jax.pmap( |
| pmapped_step, |
| axis_name="devices", |
| devices=devices, |
| ) |
|
|
| pmapped_cache = {} |
|
|
| def get_pmapped_step(length): |
| if length not in pmapped_cache: |
| pmapped_cache[length] = make_pmapped_step(length) |
| return pmapped_cache[length] |
|
|
| def train_step(step, current_params, current_opt_state, current_ema_params, batch_target, batch_forcing): |
| length = schedule_length(step) |
| batch_target = _slice_time(batch_target, length) |
| batch_forcing = _slice_time(batch_forcing, length) |
| pmapped = get_pmapped_step(length) |
| keys = step_rngs(step, global_batch) |
| keys = keys.reshape((requested_devices, local_batch) + keys.shape[1:]) |
| sharded_target = _put_batch_sharded(batch_target, devices, local_batch) |
| sharded_forcing = _put_batch_sharded(batch_forcing, devices, local_batch) |
| new_params, new_state, new_ema_params, loss, grad_finite = pmapped( |
| current_params, current_opt_state, current_ema_params, keys, |
| sharded_target, sharded_forcing, |
| ) |
| loss_host = np.asarray(jax.device_get(loss)) |
| finite_host = np.asarray(jax.device_get(grad_finite)) |
| if not np.all(np.isfinite(loss_host)) or not np.all(finite_host): |
| raise FloatingPointError( |
| "NeuralGCM loss became NaN/Inf on one or more devices; " |
| "reduce learning_rate or use a longer/physical trajectory." |
| ) |
| return new_params, new_state, new_ema_params, loss |
|
|
| nsteps = int( |
| max_steps |
| if max_steps is not None |
| else train_cfg.get("max_steps", 1) |
| ) |
| if nsteps <= 0: |
| raise ValueError("training.max_steps must be a positive integer") |
| if start_step > nsteps: |
| raise ValueError( |
| f"Resume checkpoint is already at step {start_step}, beyond " |
| f"requested max_steps={nsteps}" |
| ) |
| checkpoint_interval = int( |
| checkpoint_interval_requested |
| if checkpoint_interval_requested is not None |
| else train_cfg.get("checkpoint_interval", 0) |
| ) |
| if checkpoint_interval < 0: |
| raise ValueError("training.checkpoint_interval must be >= 0") |
| output = ( |
| resolve_path(checkpoint_output) |
| if checkpoint_output |
| else resolve_path(config["paths"].get("checkpoint_dir", "data/checkpoint")) |
| / "model_bak.pkl" |
| ) |
| aux_ds = ds[["geopotential_at_surface", "land_sea_mask"]] |
| if "time" in aux_ds.dims: |
| aux_ds = aux_ds.isel(time=0, drop=True) |
| if "level" not in aux_ds.coords: |
| aux_ds = aux_ds.assign_coords( |
| level=np.asarray(model.data_coords.vertical.centers) |
| ) |
|
|
| def save_training_checkpoint(completed_steps: int) -> None: |
| raw_params = ( |
| train_params |
| if requested_devices == 1 |
| else _unreplicate_tree(train_params) |
| ) |
| raw_opt_state = ( |
| train_opt_state |
| if requested_devices == 1 |
| else _unreplicate_tree(train_opt_state) |
| ) |
| ema_params = ( |
| train_ema_params |
| if requested_devices == 1 |
| else _unreplicate_tree(train_ema_params) |
| ) |
| raw_params, raw_opt_state, ema_params = jax.device_get( |
| (raw_params, raw_opt_state, ema_params) |
| ) |
| inference_params = ema_params if ema_num_steps > 0 else raw_params |
| checkpoint_parameter_summary = parameter_summary(inference_params) |
| training_state = { |
| "format_version": 1, |
| "step": completed_steps, |
| "train_params": raw_params, |
| "ema_params": ema_params, |
| "opt_state": raw_opt_state, |
| "data_stream_state": prefetcher.resume_state(), |
| "contract": resume_contract, |
| } |
| save_official_checkpoint( |
| output, |
| inference_params, |
| aux_ds, |
| gin_text, |
| metadata={ |
| "mode": mode, |
| "training_steps": completed_steps, |
| "finetune_source": finetune, |
| "resume_source": str(resume_path) if resume_path else None, |
| "ema_num_steps": ema_num_steps, |
| "loss_backend": loss_backend, |
| "ensemble_size": ensemble_size, |
| "parameter_count": checkpoint_parameter_summary["count"], |
| "parameter_bytes": checkpoint_parameter_summary["nbytes"], |
| "paper_defaults": paper_defaults, |
| "training_state": training_state, |
| }, |
| ) |
| print(f"Saved resumable official-format checkpoint at step={completed_steps}: {output}") |
|
|
| crps_weight_mode = ( |
| "uniform" if crps_training and loss_config.get("variable_weights") is None |
| else "configured" if crps_training |
| else "n/a" |
| ) |
| print( |
| f"Training devices={requested_devices}/{len(available_devices)}, " |
| f"global_batch={global_batch}, local_batch={global_batch // requested_devices}, " |
| f"inner_steps={inner_steps} (data interval={data_interval}, " |
| f"model timestep={model_timestep}), loss={loss_backend}, " |
| f"loss_normalization={'explicit_weights' if loss_config.get('variable_weights') is not None else 'configured_scales'}, " |
| f"ensemble={ensemble_size}, crps_weights={crps_weight_mode}, " |
| f"lr_peak={effective_lr:g}, configured_long_run={paper_defaults}, " |
| f"start_step={start_step}" |
| ) |
| completed_steps = start_step |
| last_saved_step = None |
| try: |
| for step in range(start_step, nsteps): |
| total_start = time.perf_counter() |
| _, samples = prefetcher.next_batch() |
| batch_target, batch_forcing = convert_samples(samples) |
| train_params, train_opt_state, train_ema_params, loss = train_step( |
| step, train_params, train_opt_state, train_ema_params, batch_target, batch_forcing |
| ) |
| loss_value = float(np.asarray(jax.device_get(loss)).reshape(-1)[0]) |
| elapsed = time.perf_counter() - total_start |
| throughput = global_batch / elapsed |
| current_lr = float(np.asarray(jax.device_get(schedule(step)))) |
| print( |
| f"step={step + 1}/{nsteps} loss={loss_value:.6g} " |
| f"lr={current_lr:.6g} rollout_hours={(schedule_length(step) - 1) * int(config['data'].get('time_step_hours', 6))} " |
| f"elapsed={elapsed:.2f}s throughput={throughput:.4f} samples/s" |
| ) |
| completed_steps = step + 1 |
| if checkpoint_interval and completed_steps % checkpoint_interval == 0: |
| save_training_checkpoint(completed_steps) |
| last_saved_step = completed_steps |
| finally: |
| prefetcher.close() |
| if last_saved_step != completed_steps: |
| save_training_checkpoint(completed_steps) |
|
|
|
|
| def main(*, forced_mode: str | None = None) -> None: |
| parser = argparse.ArgumentParser() |
| parser.add_argument("--config", default="conf/config.yaml") |
| parser.add_argument("--mode") |
| parser.add_argument( |
| "--finetune", |
| nargs="?", |
| const="official", |
| help="load params only and reset optimizer (no value selects this mode's official checkpoint)", |
| ) |
| parser.add_argument( |
| "--resume", |
| nargs="?", |
| const="default", |
| help="restore params, EMA, optimizer, step and data stream (no value selects the output checkpoint)", |
| ) |
| parser.add_argument("--data-dir") |
| parser.add_argument("--max-steps", type=int) |
| parser.add_argument("--learning-rate", type=float) |
| parser.add_argument( |
| "--paper-defaults", |
| action="store_true", |
| help="use configured long-run settings informed by the public training description", |
| ) |
| parser.add_argument( |
| "--trajectory-length", |
| type=int, |
| help="override the configured rollout curriculum (includes the initial frame)", |
| ) |
| parser.add_argument( |
| "--checkpoint-output", |
| help="explicit output checkpoint path (default: paths.checkpoint_dir/model_bak.pkl)", |
| ) |
| parser.add_argument( |
| "--checkpoint-interval", |
| type=int, |
| help="save resumable state every N completed steps; 0 saves only at exit", |
| ) |
| parser.add_argument( |
| "--loss-backend", |
| choices=("official", "paper", "legacy_official", "scaled", "crps"), |
| help=( |
| "override training.loss.backend (official/paper use the published " |
| "five-term deterministic objective; scaled is for synthetic smoke data)" |
| ), |
| ) |
| parser.add_argument( |
| "--devices", |
| type=int, |
| help="number of local JAX devices for synchronous data parallel training", |
| ) |
| parser.add_argument( |
| "--data-workers", |
| type=int, |
| help="host threads used to prefetch OneScience ERA5Dataset windows", |
| ) |
| parser.add_argument( |
| "--prefetch-batches", |
| type=int, |
| help="number of full host batches queued ahead of the train step", |
| ) |
| parser.add_argument("--validate-only", action="store_true") |
| 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_mode = args.mode or config["training"].get("mode", "weather_forecast") |
| requested_mode = MODE_ALIASES.get(requested_mode, requested_mode) |
| if forced_mode is not None: |
| fixed_mode = MODE_ALIASES.get(forced_mode, forced_mode) |
| if args.mode is not None and requested_mode != fixed_mode: |
| raise ValueError( |
| f"This launcher is fixed to mode={fixed_mode!r}; received " |
| f"conflicting --mode {args.mode!r}. Use scripts/train.py to " |
| "select a mode dynamically." |
| ) |
| mode = fixed_mode |
| else: |
| mode = requested_mode |
| if mode not in config["model"].get("profiles", {}): |
| raise ValueError(f"Unknown NeuralGCM mode {mode!r}") |
| finetune = args.finetune |
| if finetune == "official": |
| finetune = config["model"]["profiles"][mode]["official_reference"] |
| resume = args.resume |
| if resume == "default": |
| resume = args.checkpoint_output or str( |
| resolve_path(config["paths"].get("checkpoint_dir", "data/checkpoint")) |
| / "model_bak.pkl" |
| ) |
| if args.validate_only: |
| validation_years = list(config["data"].get("train_years", [2000])) |
| load_era5_dataset(config, validation_years) |
| if era5_data_is_synthetic(config, validation_years): |
| validate_synthetic_era5_version(config, validation_years) |
| print("OneScience ERA5Dataset validation complete") |
| return |
| train( |
| config, |
| mode, |
| finetune, |
| args.max_steps, |
| args.learning_rate, |
| args.devices, |
| args.data_workers, |
| args.prefetch_batches, |
| args.trajectory_length, |
| args.checkpoint_output, |
| args.paper_defaults, |
| resume, |
| args.checkpoint_interval, |
| args.loss_backend, |
| ) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|