NeuralGCM / scripts /losses.py
yzt15806542928's picture
Upload folder using huggingface_hub
f4a39ee verified
Raw
History Blame Contribute Delete
26.1 kB
"""Official-style NeuralGCM trajectory losses.
The released NeuralGCM repository contains the public metric primitives, but
the Gin binding used by Google's training jobs is proprietary. This module
assembles the deterministic five-term objective described in Supplementary
G.3/G.4 while exposing unpublished numerical tables explicitly in YAML. Both
the deterministic loss and the public two-member CRPS objective are built from
the released implementations.
"""
from __future__ import annotations
from collections.abc import Mapping
import functools
from typing import Any
import jax
import jax.numpy as jnp
import numpy as np
def _leaf_items(tree: Any, prefix: tuple[str, ...] = ()):
if isinstance(tree, Mapping):
for key, value in tree.items():
yield from _leaf_items(value, prefix + (str(key),))
else:
yield prefix, tree
def _lookup_scale(path: tuple[str, ...], scales: Mapping[str, float]) -> float:
"""Returns a physical-unit scale using the leaf name as the key."""
leaf = path[-1] if path else "default"
# A tracer's full path is checked first, then its leaf name, then default.
full = ".".join(path)
value = scales.get(full, scales.get(leaf, scales.get("default", 1.0)))
return max(float(value), 1e-12)
def _canonical_variable(path: tuple[str, ...]) -> str:
"""Map pressure-level and model-state names to configured loss groups."""
leaf = path[-1] if path else "default"
aliases = {
"geopotential": "z",
"temperature": "t",
"temperature_variation": "t",
"u_component_of_wind": "u",
"v_component_of_wind": "v",
}
return aliases.get(leaf, leaf)
def _lookup_named_value(
path: tuple[str, ...], values: Mapping[str, Any], default: float
) -> Any:
full = ".".join(path)
leaf = path[-1] if path else "default"
canonical = _canonical_variable(path)
return values.get(
full,
values.get(leaf, values.get(canonical, values.get("default", default))),
)
def _broadcast_level_value(value: Any, error, *, name: str):
"""Broadcast a scalar or pressure/sigma-level vector over a trajectory."""
value = jnp.asarray(value, dtype=jnp.asarray(error).real.dtype)
if value.ndim == 0:
return value
if value.ndim != 1 or getattr(error, "ndim", 0) < 2:
raise ValueError(f"loss {name} must be scalar or a one-dimensional level vector")
if value.shape[0] != error.shape[1]:
raise ValueError(
f"loss {name} has {value.shape[0]} levels, but the trajectory has "
f"{error.shape[1]}"
)
return value.reshape((1, value.shape[0]) + (1,) * (error.ndim - 2))
def _map_named(tree: Any, fn, prefix: tuple[str, ...] = ()):
if isinstance(tree, Mapping):
return {
key: _map_named(value, fn, prefix + (str(key),))
for key, value in tree.items()
}
return fn(prefix, tree)
class _PaperVariableRescaling:
"""Paper G.3 scaling with YAML-overridable 24-hour difference scales."""
def __init__(
self,
trajectory_spec,
*,
scales: Mapping[str, Any],
factors: Mapping[str, Any],
weights: Mapping[str, Any] | None = None,
):
del trajectory_spec
self.scales = scales
self.factors = factors
self.weights = weights
def __call__(self, errors, targets):
del targets
def rescale(path, error):
if self.weights is not None:
weight = _broadcast_level_value(
_lookup_named_value(path, self.weights, 1.0),
error,
name=f"variable_weights.{'.'.join(path)}",
)
weight = jnp.maximum(weight, 0.0)
return error * jnp.sqrt(weight)
scale = _broadcast_level_value(
_lookup_named_value(path, self.scales, 1.0),
error,
name=f"variable_scales.{'.'.join(path)}",
)
scale = jnp.maximum(scale, 1e-12)
factor = _broadcast_level_value(
_lookup_named_value(path, self.factors, 1.0),
error,
name=f"variable_factors.{'.'.join(path)}",
)
return error * (factor / scale)
return _map_named(errors, rescale)
def _filter_group(path: tuple[str, ...]) -> str:
variable = _canonical_variable(path)
if variable in {
"specific_humidity",
"specific_cloud_ice_water_content",
"specific_cloud_liquid_water_content",
}:
return "moisture"
if variable in {"divergence", "vorticity", "log_surface_pressure"}:
return "divergence"
if variable in {"u", "v"}:
return "wind"
if variable == "t":
return "temperature"
return "default"
class _PaperPredictabilityFilter:
"""Order-12 lead-time filter reconstructed from Supplementary Fig. 8."""
def __init__(
self,
trajectory_spec,
*,
schedules: Mapping[str, list[float]],
lead_hours: list[float],
order: int = 12,
is_encoded: bool = False,
):
from dinosaur import filtering
self._filtering = filtering
self.grid = (
trajectory_spec.coords.horizontal
if is_encoded
else trajectory_spec.data_coords.horizontal
)
self.order = int(order)
if self.order <= 0:
raise ValueError("loss.predictability_filter.order must be positive")
n = int(trajectory_spec.trajectory_length)
source_hours = np.asarray(lead_hours, dtype=np.float64)
if source_hours.ndim != 1 or source_hours.size == 0:
raise ValueError("loss.predictability_filter.lead_hours must be non-empty")
if np.any(np.diff(source_hours) <= 0):
raise ValueError("loss.predictability_filter.lead_hours must increase")
target_hours = np.arange(n, dtype=np.float64) * float(
trajectory_spec.steps_per_save
)
self.cutoffs = {}
for group, values in schedules.items():
values = np.asarray(values, dtype=np.float64)
if values.shape != source_hours.shape:
raise ValueError(
f"loss.predictability_filter.cutoffs.{group} has "
f"{values.size} entries; expected {source_hours.size}"
)
self.cutoffs[str(group)] = np.interp(
target_hours, source_hours, values
)
if "default" not in self.cutoffs:
raise ValueError("loss.predictability_filter.cutoffs.default is required")
def __call__(self, errors, targets):
del targets
max_wavenumber = float(np.max(np.asarray(self.grid.modal_axes[1])))
def apply_filter(path, error):
if getattr(error, "ndim", 0) < 2:
return error
group = _filter_group(path)
cutoffs = self.cutoffs.get(group, self.cutoffs["default"])
cutoffs = np.clip(cutoffs, 1.0, max_wavenumber)
# dinosaur.exponential_filter uses normalized total wavenumber.
# Choose attenuation so the response is 0.5 at each configured
# absolute cutoff. This preserves the paper's order-12 profile.
attenuation = np.log(2.0) * np.power(
max_wavenumber / cutoffs, 2 * self.order
)
attenuation = attenuation.reshape((-1,) + (1,) * (error.ndim - 1))
filter_fn = self._filtering.exponential_filter(
self.grid,
attenuation=jnp.asarray(attenuation, dtype=jnp.float32),
order=self.order,
)
return filter_fn(error)
return _map_named(errors, apply_filter)
class _PaperDeterministicLoss:
"""Five-term deterministic objective from Supplementary section G.4."""
def __init__(self, terms, bias_metric, coefficients: Mapping[str, float]):
self.terms = terms
self.bias_metric = bias_metric
self.coefficients = coefficients
@staticmethod
def _global_bias_per_example(metric, prediction, target, axis_names):
"""Extend the released bias metric over local and device batch axes."""
from model.reference_code import linear_transforms
prediction = metric.get_representation(prediction)
target = metric.get_representation(target)
truncate = metric.transform.transforms[0]
if not isinstance(truncate, linear_transforms.TruncateToTrajectoryLength):
raise TypeError("BatchMeanSquaredBias must start with trajectory truncation")
prediction = metric.getter(truncate(prediction, None))
target = metric.getter(truncate(target, None))
prediction = metric.metric_fn(prediction)
target = metric.metric_fn(target)
prediction = jax.tree_util.tree_map(
lambda value: jax.lax.pmean(value, axis_name=axis_names), prediction
)
target = jax.tree_util.tree_map(
lambda value: jax.lax.pmean(value, axis_name=axis_names), target
)
prediction = jax.tree_util.tree_map(
lambda value: jnp.mean(value, axis=0, keepdims=True), prediction
)
target = jax.tree_util.tree_map(
lambda value: jnp.mean(value, axis=0, keepdims=True), target
)
errors = jax.tree_util.tree_map(jnp.subtract, prediction, target)
errors = metric.transform(errors, target)
per_variable = jax.tree_util.tree_map(
lambda value: jnp.mean(jnp.square(value)), errors
)
return sum(jax.tree_util.tree_leaves(per_variable))
def evaluate_batch(self, prediction, target, *, device_axis_name=None):
"""Evaluate one global batch, including a true global spectral bias."""
values = {}
for name, metric in self.terms.items():
per_example = jax.vmap(metric.evaluate, in_axes=(0, 0))(
prediction, target
)
values[name] = jnp.mean(per_example)
axis_names = (
("loss_batch",)
if device_axis_name is None
else ("loss_batch", device_axis_name)
)
bias = jax.vmap(
functools.partial(
self._global_bias_per_example,
self.bias_metric,
axis_names=axis_names,
),
in_axes=(0, 0),
axis_name="loss_batch",
)(prediction, target)
values["bias"] = jnp.mean(bias)
return sum(
self.coefficients[name] * value for name, value in values.items()
)
def __call__(self, prediction, target):
prediction = jax.tree_util.tree_map(lambda value: value[None], prediction)
target = jax.tree_util.tree_map(lambda value: value[None], target)
return self.evaluate_batch(prediction, target)
def _time_factor(n_time: int, steps_per_save: int, mode: str) -> jnp.ndarray:
"""Public NeuralGCM time rescaling, returned as squared-error factors."""
if n_time <= 0:
return jnp.ones((0,), dtype=jnp.float32)
if mode == "none":
return jnp.ones((n_time,), dtype=jnp.float32)
if mode == "legacy":
# linear_transforms.LegacyTimeRescaling: errors /
# sqrt((trajectory_length - 1) * steps_per_save).
denominator = max((n_time - 1) * int(steps_per_save), 1)
return jnp.full((n_time,), 1.0 / denominator, dtype=jnp.float32)
if mode == "random_walk":
# Same normalized inverse-variance weighting as the public
# TimeRescaling transform with base_squared_error_in_hours=1.
t = jnp.arange(n_time, dtype=jnp.float32) * float(steps_per_save)
inv_variance = 1.0 / (1.0 + t)
return inv_variance / jnp.sum(inv_variance)
raise ValueError(f"Unknown loss.time_rescaling mode: {mode!r}")
def _surface_mean(square_error, coords):
"""Computes the public metrics_util.nodal_surface_mean for one leaf."""
horizontal = coords.horizontal
expected = tuple(horizontal.nodal_shape[-2:])
if getattr(square_error, "ndim", 0) >= 2 and tuple(square_error.shape[-2:]) == expected:
surface_area = 4 * jnp.pi * horizontal.radius**2
return horizontal.integrate(square_error) / surface_area
# Metadata or scalar diagnostics are not part of the official loss, but
# retaining a finite scalar here makes custom profiles easier to inspect.
return jnp.mean(square_error)
def _per_leaf_loss(error, path, coords, scales, level_weights, time_factors):
error = jnp.asarray(error)
if error.dtype.kind in ("O", "U", "S"):
return jnp.asarray(0.0, dtype=jnp.float32)
# Trajectory representations conventionally use [time, level, lon, lat].
# A few surface fields omit the level axis; both are handled by broadcasting
# the time factor along all remaining dimensions.
if error.ndim == 0:
return jnp.mean(jnp.square(error / _lookup_scale(path, scales)))
factor = time_factors
if error.shape[0] != factor.shape[0]:
# The prediction may omit the initialization frame. The caller aligns
# target and prediction; this defensive slice handles custom adapters.
factor = factor[-error.shape[0]:]
reshape = (factor.shape[0],) + (1,) * (error.ndim - 1)
transformed = error / _lookup_scale(path, scales)
transformed = transformed * jnp.sqrt(factor.reshape(reshape))
squared = jnp.square(transformed)
if level_weights and squared.ndim >= 4:
weights = jnp.asarray(level_weights, dtype=squared.dtype)
weights = weights[: squared.shape[1]]
squared = squared * weights.reshape((1, weights.shape[0]) + (1,) * (squared.ndim - 2))
return jnp.mean(_surface_mean(squared, coords))
def make_loss_fn(
model,
*,
steps_per_save: int,
trajectory_length: int | None = None,
config: Mapping[str, Any] | None = None,
mode: str | None = None,
):
"""Build a JAX-compatible trajectory loss.
``backend=official`` uses the paper's five-term deterministic objective,
assembled from the released metric primitives. ``backend=legacy_official``
retains the earlier public WeightedL2CumulativeLoss baseline.
``backend=crps`` builds the released two-member nodal + spectral CRPS
objective described in supplementary section G.6. ``backend=scaled``
keeps the historical deterministic approximation.
"""
cfg = dict(config or {})
backend = str(cfg.get("backend", "official")).lower()
if backend in {"official", "paper", "legacy_official", "crps"}:
if trajectory_length is None:
raise ValueError(f"trajectory_length is required for the {backend} loss backend")
from model.reference_code import metrics_util
trajectory_spec = metrics_util.TrajectorySpec(
trajectory_length=int(trajectory_length),
max_trajectory_length=int(trajectory_length),
steps_per_save=int(steps_per_save),
coords=model.coords,
data_coords=model.data_coords,
)
if backend == "crps":
from model.reference_code import linear_transforms
from model.reference_code import stochastic_losses
weights = cfg.get("variable_weights")
variable_scale = float(cfg.get("variable_scale", 1.0))
nodal_hours = float(cfg.get("nodal_time_scale_hours", 24.0))
spectral_hours = float(cfg.get("spectral_time_scale_hours", 40.0))
max_wavenumber = int(cfg.get("spectral_max_wavenumber", 80))
if nodal_hours <= 0 or spectral_hours <= 0:
raise ValueError("CRPS time scale hours must be positive")
if max_wavenumber <= 0:
raise ValueError("loss.spectral_max_wavenumber must be positive")
variable_rescaling = functools.partial(
linear_transforms.PerVariableRescaling,
weights=weights,
scale=variable_scale,
)
nodal_time_rescaling = functools.partial(
linear_transforms.DelayedTimeRescaling,
base_squared_error_in_hours=nodal_hours,
delay_power=1.0,
decay_power=1.0,
)
spectral_time_rescaling = functools.partial(
linear_transforms.DelayedTimeRescaling,
base_squared_error_in_hours=spectral_hours,
delay_power=4.0,
decay_power=1.0,
)
wavenumber_mask = functools.partial(
linear_transforms.TotalWavenumberMasking,
max_wavenumber=max_wavenumber,
is_encoded=False,
)
nodal_crps = stochastic_losses.CRPSLoss(
trajectory_spec,
components=(variable_rescaling, nodal_time_rescaling),
beta=1.0,
ensemble_term_weight=0.5,
is_nodal=True,
is_encoded=False,
)
spectral_crps = stochastic_losses.CRPSLoss(
trajectory_spec,
components=(
variable_rescaling,
spectral_time_rescaling,
wavenumber_mask,
),
beta=1.0,
ensemble_term_weight=0.5,
is_nodal=False,
is_encoded=False,
)
def crps_loss(prediction, target):
return nodal_crps.evaluate(prediction, target) + spectral_crps.evaluate(
prediction, target
)
return crps_loss
from model.reference_code import linear_transforms
from model.reference_code import metrics
if backend == "legacy_official":
public_loss = metrics.WeightedL2CumulativeLoss(
trajectory_spec, weights=None, scale=float(cfg.get("scale", 1.0))
)
return public_loss.evaluate
if mode is None:
raise ValueError("mode is required for the paper deterministic loss")
cutoffs_by_mode = dict(
cfg.get(
"spectral_cutoff_by_mode",
{
"weather_forecast": 120,
"climate_scale": 80,
"forecast_2_8_deg": 42,
},
)
)
if mode not in cutoffs_by_mode:
raise ValueError(
f"No deterministic spectral cutoff configured for {mode!r}"
)
spectral_cutoff = int(cutoffs_by_mode[mode])
scales = dict(cfg.get("variable_scales", {}))
factors = dict(cfg.get("variable_factors", {}))
explicit_weights = cfg.get("variable_weights")
variable_rescaling = functools.partial(
_PaperVariableRescaling,
scales=scales,
factors=factors,
weights=None if explicit_weights is None else dict(explicit_weights),
)
accuracy_time = functools.partial(
linear_transforms.DelayedTimeRescaling,
base_squared_error_in_hours=float(
cfg.get("accuracy_time_scale_hours", 24.0)
),
delay_power=1.0,
decay_power=1.0,
)
spectral_time = functools.partial(
linear_transforms.DelayedTimeRescaling,
base_squared_error_in_hours=float(
cfg.get("spectral_time_scale_hours", 40.0)
),
delay_power=4.0,
decay_power=1.0,
)
filter_cfg = dict(cfg.get("predictability_filter", {}))
filter_enabled = bool(filter_cfg.get("enabled", True))
def accuracy_components(is_encoded: bool):
components = [variable_rescaling, accuracy_time]
if filter_enabled:
components.append(
functools.partial(
_PaperPredictabilityFilter,
schedules=dict(filter_cfg.get("cutoffs", {})),
lead_hours=list(filter_cfg.get("lead_hours", [])),
order=int(filter_cfg.get("order", 12)),
is_encoded=is_encoded,
)
)
return tuple(components)
def spectrum_components(is_encoded: bool):
return (
variable_rescaling,
spectral_time,
functools.partial(
linear_transforms.TotalWavenumberMasking,
max_wavenumber=spectral_cutoff,
is_encoded=is_encoded,
),
)
terms = {
"data": metrics.TransformedL2Loss(
trajectory_spec,
components=accuracy_components(False),
is_nodal=False,
is_encoded=False,
),
"data_spectrum": metrics.TransformedL2SpectrumLoss(
trajectory_spec,
components=spectrum_components(False),
is_nodal=False,
is_encoded=False,
),
"model": metrics.TransformedL2Loss(
trajectory_spec,
components=accuracy_components(True),
is_nodal=False,
is_encoded=True,
),
"model_spectrum": metrics.TransformedL2SpectrumLoss(
trajectory_spec,
components=spectrum_components(True),
is_nodal=False,
is_encoded=True,
),
}
bias_metric = metrics.BatchMeanSquaredBias(
trajectory_spec,
components=(variable_rescaling,),
is_nodal=False,
is_encoded=False,
)
coefficients = {
"data": float(cfg.get("data_weight", 20.0)),
"data_spectrum": float(cfg.get("data_spectrum_weight", 0.1)),
"model": float(cfg.get("model_weight", 1.0)),
"model_spectrum": float(cfg.get("model_spectrum_weight", 0.1)),
"bias": float(cfg.get("bias_weight", 2.0)),
}
return _PaperDeterministicLoss(terms, bias_metric, coefficients)
if backend not in {"scaled", "legacy"}:
raise ValueError(f"Unknown loss.backend {backend!r}")
scales = dict(cfg.get("variable_scales", {}))
scales.setdefault("z", 1.0e4) # geopotential (m² s⁻²)
scales.setdefault("t", 30.0) # temperature (K)
scales.setdefault("u", 30.0) # zonal wind (m s⁻¹)
scales.setdefault("v", 30.0) # meridional wind (m s⁻¹)
scales.setdefault("specific_humidity", 1.0e-2)
scales.setdefault("default", 1.0)
level_weights = cfg.get("level_weights", [])
time_mode = str(cfg.get("time_rescaling", "legacy"))
spectral_weight = float(cfg.get("spectral_weight", 0.0))
bias_weight = float(cfg.get("bias_weight", 0.0))
coords = model.data_coords
def loss_fn(prediction, target):
pred = dict(prediction.data_nodal_trajectory)
truth = dict(target.data_nodal_trajectory)
pred.pop("sim_time", None)
truth.pop("sim_time", None)
def align(a, b):
if getattr(a, "ndim", 0) and getattr(b, "ndim", 0):
if a.shape[0] != b.shape[0] and a.shape[1:] == b.shape[1:]:
return b[-a.shape[0]:]
return b
truth = __import__("jax").tree_util.tree_map(align, pred, truth)
leaves = []
first_array = next(
(value for _, value in _leaf_items(pred) if getattr(value, "ndim", 0)),
None,
)
if first_array is None:
return jnp.asarray(0.0, dtype=jnp.float32)
factors = _time_factor(int(first_array.shape[0]), steps_per_save, time_mode)
for path, p in _leaf_items(pred):
# Resolve the corresponding target leaf without assuming a flat tree.
node = truth
for key in path:
node = node[key]
if getattr(p, "dtype", None) is None or p.dtype.kind in ("O", "U", "S"):
continue
leaves.append(_per_leaf_loss(p - node, path, coords, scales, level_weights, factors))
accuracy = jnp.sum(jnp.stack(leaves)) if leaves else jnp.asarray(0.0)
# Optional public spectral-norm term. It is disabled by default because
# the private Gin files do not expose its coefficient.
spectral = jnp.asarray(0.0)
if spectral_weight:
pmodal = dict(prediction.data_modal_trajectory)
tmodal = dict(target.data_modal_trajectory)
pmodal.pop("sim_time", None)
tmodal.pop("sim_time", None)
tmodal = __import__("jax").tree_util.tree_map(align, pmodal, tmodal)
terms = []
for path, p in _leaf_items(pmodal):
node = tmodal
for key in path:
node = node[key]
if getattr(p, "ndim", 0) >= 4:
# Public _compute_spectral_norm: norm over longitude
# wavenumber, followed by MSE and variable scaling.
ps = jnp.sqrt(jnp.sum(jnp.real(p * jnp.conj(p)), axis=-2, keepdims=True) + 1e-12)
ts = jnp.sqrt(jnp.sum(jnp.real(node * jnp.conj(node)), axis=-2, keepdims=True) + 1e-12)
terms.append(jnp.mean(jnp.square((ps - ts) / _lookup_scale(path, scales))))
if terms:
spectral = jnp.sum(jnp.stack(terms))
bias = jnp.asarray(0.0)
if bias_weight:
for path, p in _leaf_items(pred):
node = truth
for key in path:
node = node[key]
pmean = jnp.mean(p, axis=0)
tmean = jnp.mean(node, axis=0)
bias = bias + jnp.mean(jnp.square((pmean - tmean) / _lookup_scale(path, scales)))
return accuracy + spectral_weight * spectral + bias_weight * bias
return loss_fn