ScaleAdaptiveCM / model /scale_adaptive_cm.py
zhangrenchao's picture
Publish ScaleAdaptiveCM reproduction
90cbf31 verified
Raw
History Blame Contribute Delete
3.75 kB
"""Compact consistency model for full-grid precipitation downscaling."""
import json
from pathlib import Path
import numpy as np
import torch
from torch import nn
import torch.nn.functional as F
import yaml
def load_config(root):
return yaml.safe_load((Path(root) / "conf/config.yaml").read_text())
class TimeBlock(nn.Module):
def __init__(self, cin, cout, time_dim):
super().__init__()
groups = min(4, cout)
self.conv = nn.Conv2d(cin, cout, 3, padding=1)
self.norm = nn.GroupNorm(groups, cout)
self.time = nn.Linear(time_dim, cout)
def forward(self, x, embedding):
return F.silu(self.norm(self.conv(x)) + self.time(embedding)[:, :, None, None])
class ScaleAdaptiveCM(nn.Module):
def __init__(self, channels=(4, 8, 16), time_dim=16, sigma_data=0.5):
super().__init__()
self.sigma_data = float(sigma_data)
self.time_dim = int(time_dim)
self.time_mlp = nn.Sequential(nn.Linear(time_dim, time_dim), nn.SiLU(), nn.Linear(time_dim, time_dim))
self.enc1 = TimeBlock(1, channels[0], time_dim)
self.enc2 = TimeBlock(channels[0], channels[1], time_dim)
self.mid = TimeBlock(channels[1], channels[2], time_dim)
self.dec2 = TimeBlock(channels[2] + channels[1], channels[1], time_dim)
self.dec1 = TimeBlock(channels[1] + channels[0], channels[0], time_dim)
self.out = nn.Conv2d(channels[0], 1, 1)
self.model_config = {"channels": list(channels), "time_dim": time_dim, "sigma_data": sigma_data}
def embed_time(self, t):
half = self.time_dim // 2
freq = torch.exp(torch.linspace(0, -7, half, device=t.device))
emb = torch.cat((torch.sin(t[:, None] * freq), torch.cos(t[:, None] * freq)), 1)
return self.time_mlp(emb)
def forward(self, noisy, t):
if noisy.ndim != 4 or noisy.shape[1] != 1:
raise ValueError("expected [B,1,H,W]")
emb = self.embed_time(t.float())
e1 = self.enc1(noisy, emb)
e2 = self.enc2(F.avg_pool2d(e1, 2), emb)
mid = self.mid(F.avg_pool2d(e2, 2), emb)
d2 = self.dec2(torch.cat((F.interpolate(mid, e2.shape[-2:], mode="bilinear", align_corners=False), e2), 1), emb)
d1 = self.dec1(torch.cat((F.interpolate(d2, e1.shape[-2:], mode="bilinear", align_corners=False), e1), 1), emb)
raw = self.out(d1)
sigma2 = self.sigma_data ** 2
cskip = sigma2 / ((t[:, None, None, None] - 0.002).square() + sigma2)
cout = self.sigma_data * t[:, None, None, None] / torch.sqrt(t[:, None, None, None].square() + sigma2)
return cskip * noisy + cout * raw
def structured_fields(samples, high_h, high_w, seed):
rng = np.random.default_rng(seed)
yy, xx = np.mgrid[-1:1:complex(high_h), -1:1:complex(high_w)]
fields = []
for i in range(samples):
phase = 2 * np.pi * i / max(samples, 4)
itcz = 9 * np.exp(-((yy - .12 * np.sin(phase)) / .16) ** 2)
storms = 18 * np.exp(-((xx - .45 * np.cos(phase)) ** 2 + (yy - .3 * np.sin(phase)) ** 2) / .025)
texture = 2 * np.maximum(0, np.sin(18 * xx + phase) * np.cos(13 * yy - phase))
fields.append(np.maximum(0, itcz + storms + texture + rng.normal(0, .15, yy.shape)))
return np.asarray(fields, np.float32)[:, None]
def radial_spectrum(field):
power = np.abs(np.fft.fftshift(np.fft.fft2(field))) ** 2
y, x = np.indices(field.shape); r = np.sqrt((y-field.shape[0]/2)**2 + (x-field.shape[1]/2)**2).astype(int)
return np.bincount(r.ravel(), power.ravel()) / np.maximum(np.bincount(r.ravel()), 1)
def write_json(path, value):
path = Path(path); path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(json.dumps(value, indent=2) + "\n")