"""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