NNCAM / model /nncam.py
zhangrenchao's picture
Publish NNCAM reproduction
8791cfe verified
Raw History Blame Contribute Delete
4.47 kB
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