File size: 3,752 Bytes
90cbf31
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
"""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")