Ai2_Climate_Emulator / model /normalization.py
yzt15806542928's picture
Upload folder using huggingface_hub
5c365c5 verified
Raw
History Blame Contribute Delete
6.75 kB
"""Full-field and residual normalization from Appendix H."""
from __future__ import annotations
from pathlib import Path
import numpy as np
class ResidualNormalizer:
def __init__(self, eps: float = 1e-6) -> None:
self.eps = float(eps)
self.mean: np.ndarray | None = None
self.std: np.ndarray | None = None
self.residual_scale: np.ndarray | None = None
def fit(self, current: np.ndarray, next_values: np.ndarray | None = None) -> "ResidualNormalizer":
cur = np.asarray(current, dtype=np.float64)
if cur.ndim != 4:
raise ValueError("normalizer expects [N,C,H,W]")
self.mean = cur.mean(axis=(0, 2, 3))
self.std = np.maximum(cur.std(axis=(0, 2, 3)), self.eps)
if next_values is None:
self.residual_scale = np.ones_like(self.std)
else:
nxt = np.asarray(next_values, dtype=np.float64)
if nxt.shape != cur.shape:
raise ValueError("current and next_values must have equal shapes")
shape = (1, -1, 1, 1)
ff_cur = (cur - self.mean.reshape(shape)) / self.std.reshape(shape)
ff_next = (nxt - self.mean.reshape(shape)) / self.std.reshape(shape)
increment_std = np.maximum((ff_next - ff_cur).std(axis=(0, 2, 3)), self.eps)
geometric_mean = float(np.exp(np.mean(np.log(increment_std))))
self.residual_scale = np.maximum(increment_std / max(geometric_mean, self.eps), self.eps)
return self
def _stats(self) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
if self.mean is None or self.std is None or self.residual_scale is None:
raise RuntimeError("normalizer has not been fitted")
return self.mean, self.std, self.residual_scale
def transform(self, values: np.ndarray) -> np.ndarray:
mean, std, scale = self._stats()
shape = (1, -1, 1, 1)
return ((np.asarray(values) - mean.reshape(shape)) / std.reshape(shape)) / scale.reshape(shape)
def inverse_transform(self, values: np.ndarray) -> np.ndarray:
mean, std, scale = self._stats()
shape = (1, -1, 1, 1)
return np.asarray(values) * scale.reshape(shape) * std.reshape(shape) + mean.reshape(shape)
def save(self, path: str | Path) -> None:
mean, std, scale = self._stats()
np.savez(path, mean=mean, std=std, residual_scale=scale, eps=np.asarray(self.eps))
@classmethod
def load(cls, path: str | Path) -> "ResidualNormalizer":
data = np.load(path)
obj = cls(float(data["eps"]) if "eps" in data else 1e-6)
obj.mean, obj.std, obj.residual_scale = data["mean"], data["std"], data["residual_scale"]
return obj
class ACEDataNormalizer:
"""Residual-scaling normalizer for ACE's 40-input/44-output contract.
Prognostic output increments use the same per-variable statistics as the
input prognostic state. Forcing and diagnostic channels use full-field
statistics. Statistics are fitted on the training split and serialized in
the checkpoint so inference uses exactly the same transform.
"""
def __init__(self, eps: float = 1e-6) -> None:
self.eps = float(eps)
self.input_mean: np.ndarray | None = None
self.input_std: np.ndarray | None = None
self.target_mean: np.ndarray | None = None
self.target_std: np.ndarray | None = None
self.residual_scale: np.ndarray | None = None
def fit(self, inputs: np.ndarray, targets: np.ndarray) -> "ACEDataNormalizer":
x = np.asarray(inputs, dtype=np.float64)
y = np.asarray(targets, dtype=np.float64)
if x.ndim != 4 or y.ndim != 4 or x.shape[0] != y.shape[0]:
raise ValueError("inputs and targets must be [N,C,H,W] with equal N")
self.input_mean = x.mean(axis=(0, 2, 3))
self.input_std = np.maximum(x.std(axis=(0, 2, 3)), self.eps)
self.target_mean = y.mean(axis=(0, 2, 3))
self.target_std = np.maximum(y.std(axis=(0, 2, 3)), self.eps)
prognostic = x.shape[1] - 6
x_state = (x[:, :prognostic] - self.input_mean[:prognostic].reshape(1, -1, 1, 1)) / self.input_std[:prognostic].reshape(1, -1, 1, 1)
y_state = (y[:, :prognostic] - self.target_mean[:prognostic].reshape(1, -1, 1, 1)) / self.target_std[:prognostic].reshape(1, -1, 1, 1)
increment_std = np.maximum((y_state - x_state).std(axis=(0, 2, 3)), self.eps)
geometric_mean = float(np.exp(np.mean(np.log(increment_std))))
self.residual_scale = np.ones(y.shape[1], dtype=np.float64)
self.residual_scale[:prognostic] = np.maximum(increment_std / max(geometric_mean, self.eps), self.eps)
return self
def _stats(self) -> tuple[np.ndarray, np.ndarray, np.ndarray, np.ndarray, np.ndarray]:
if any(value is None for value in (self.input_mean, self.input_std, self.target_mean, self.target_std, self.residual_scale)):
raise RuntimeError("normalizer has not been fitted")
return self.input_mean, self.input_std, self.target_mean, self.target_std, self.residual_scale
def transform_inputs(self, inputs: np.ndarray) -> np.ndarray:
mean, std, _, _, _ = self._stats()
values = np.asarray(inputs)
shape = (1,) * (values.ndim - 3) + (-1, 1, 1)
return ((values - mean.reshape(shape)) / std.reshape(shape)).astype(np.float32)
def transform_targets(self, targets: np.ndarray) -> np.ndarray:
_, _, mean, std, scale = self._stats()
values = np.asarray(targets)
shape = (1,) * (values.ndim - 3) + (-1, 1, 1)
return (((values - mean.reshape(shape)) / std.reshape(shape)) / scale.reshape(shape)).astype(np.float32)
def inverse_targets(self, targets: np.ndarray) -> np.ndarray:
_, _, mean, std, scale = self._stats()
values = np.asarray(targets)
shape = (1,) * (values.ndim - 3) + (-1, 1, 1)
return (values * scale.reshape(shape) * std.reshape(shape) + mean.reshape(shape)).astype(np.float32)
def to_dict(self) -> dict[str, np.ndarray | float]:
mean, std, target_mean, target_std, scale = self._stats()
return {"eps": self.eps, "input_mean": mean, "input_std": std, "target_mean": target_mean, "target_std": target_std, "residual_scale": scale}
@classmethod
def from_dict(cls, values: dict) -> "ACEDataNormalizer":
obj = cls(float(values.get("eps", 1e-6)))
obj.input_mean = np.asarray(values["input_mean"])
obj.input_std = np.asarray(values["input_std"])
obj.target_mean = np.asarray(values["target_mean"])
obj.target_std = np.asarray(values["target_std"])
obj.residual_scale = np.asarray(values["residual_scale"])
return obj