import math import numpy as np import torch from torch import nn INPUT_GROUPS = {"T": slice(0, 30), "Q": slice(30, 60), "V": slice(60, 90), "Ps": slice(90, 91), "Sin": slice(91, 92), "H": slice(92, 93), "E": slice(93, 94)} OUTPUT_GROUPS = {"dT": slice(0, 30), "dQ": slice(30, 60), "SWtoa": slice(60, 61), "SWsfc": slice(61, 62), "LWtoa": slice(62, 63), "LWsfc": slice(63, 64), "P": slice(64, 65)} CP = 1004.0 LV = 2.5e6 OUTPUT_SCALE = np.r_[np.full(30, CP), np.full(30, LV), np.full(4, 1e-3), 2e-2].astype(np.float32) class NNCAM(nn.Module): """Fully connected 94-to-65 atmospheric-column parameterization.""" def __init__(self, input_dim=94, output_dim=65, width=32, depth=4, negative_slope=0.3): super().__init__() if input_dim != 94 or output_dim != 65: raise ValueError("NNCAM requires 94 inputs and 65 outputs") if width < 1 or depth < 1: raise ValueError("width and depth must be positive") layers = [] in_features = input_dim for _ in range(depth): layers.extend((nn.Linear(in_features, width), nn.LeakyReLU(negative_slope))) in_features = width layers.append(nn.Linear(in_features, output_dim)) self.network = nn.Sequential(*layers) self.model_config = {"input_dim": input_dim, "output_dim": output_dim, "width": width, "depth": depth, "negative_slope": negative_slope} def forward(self, x): if x.ndim != 2 or x.shape[1] != 94: raise ValueError(f"expected input [B, 94], got {tuple(x.shape)}") output = self.network(x) if output.shape != (x.shape[0], 65): raise RuntimeError(f"expected output [B, 65], got {tuple(output.shape)}") return output def build_model(width=32, depth=4, negative_slope=0.3): return NNCAM(width=width, depth=depth, negative_slope=negative_slope) def generate_fake_data(n_samples=192, seed=42): rng = np.random.default_rng(seed) sigma = np.linspace(0.02, 1.0, 30, dtype=np.float32)[None, :] lat = rng.uniform(-math.pi / 2, math.pi / 2, (n_samples, 1)).astype(np.float32) time = rng.uniform(0, 2 * math.pi, (n_samples, 1)).astype(np.float32) ps = rng.normal(101000.0, 1800.0, (n_samples, 1)).astype(np.float32) insolation = np.maximum(0.0, 950.0 * np.cos(lat) * (0.65 + 0.35 * np.sin(time))).astype(np.float32) height = rng.uniform(0.0, 2500.0, (n_samples, 1)).astype(np.float32) evaporation = (35.0 + 80.0 * np.maximum(np.cos(lat), 0.0) + rng.normal(0, 4, (n_samples, 1))).astype(np.float32) temperature = 205.0 + 83.0 * sigma**0.24 - 0.006 * height + 4.0 * np.cos(lat) * sigma temperature += rng.normal(0, 1.2, temperature.shape) humidity = (0.00005 + 0.017 * sigma**3 * np.maximum(np.cos(lat), 0.15)) * rng.lognormal(0, 0.12, temperature.shape) wind = 18.0 * np.sin(lat) * (1.0 - sigma) + 5.0 * np.sin(time + 3.0 * sigma) + rng.normal(0, 2, temperature.shape) instability = np.maximum(temperature[:, -1:] - temperature[:, 18:19] - 25.0, 0.0) moisture = humidity[:, -8:].mean(1, keepdims=True) precipitation = np.maximum(0.0, 4.0e-5 * instability * moisture / 0.012 + rng.normal(0, 1.5e-6, (n_samples, 1))).astype(np.float32) heating_shape = np.exp(-((sigma - 0.55) / 0.23) ** 2) drying_shape = np.exp(-((sigma - 0.78) / 0.18) ** 2) dt = (precipitation * LV / CP / 86400.0 * heating_shape - 8e-6 * (temperature - temperature.mean(1, keepdims=True))).astype(np.float32) dq = (-precipitation / 86400.0 * drying_shape + evaporation / LV / 30.0 / 86400.0).astype(np.float32) cloud = np.clip(moisture / 0.014, 0.0, 1.0) targets = (dt, dq, 0.30 * insolation, insolation * (0.72 - 0.18 * cloud), 215.0 + 0.65 * (temperature[:, -1:] - 273.0) - 24.0 * cloud, 330.0 + 1.1 * (temperature[:, -1:] - 285.0) + 18.0 * cloud, precipitation) x = np.concatenate((temperature, humidity, wind, ps, insolation, height, evaporation), axis=1).astype(np.float32) y = np.concatenate(targets, axis=1).astype(np.float32) return x, y, lat[:, 0], time[:, 0] def fit_normalizer(values): mean = values.mean(0).astype(np.float32) scale = np.maximum(np.ptp(values, axis=0), values.std(0)) return mean, np.where(scale > 1e-12, scale, 1.0).astype(np.float32) def normalize_input(values, mean, scale): return (values - mean) / scale def scale_output(values): return values * OUTPUT_SCALE def unscale_output(values): return values / OUTPUT_SCALE