File size: 6,788 Bytes
15aff58
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
"""Independent PyTorch implementation of a compact multi-branch FuXi-DA U-Net."""
import torch
from torch import nn
import torch.nn.functional as F
from torch.utils.data import Dataset

BACKGROUND_RAW_SHAPE = (70, 721, 1440)
BACKGROUND_MODEL_SHAPE = (70, 720, 1440)
AGRI_RAW_SHAPE = (8, 15, 640, 640)
ANALYSIS_SHAPE = (70, 720, 1440)
FOOTPRINT_ORIGIN = (40, 400)


def tile_origin(tile_id, tile_size=32):
    columns = 640 // tile_size
    return FOOTPRINT_ORIGIN[0] + (tile_id // columns) * tile_size, FOOTPRINT_ORIGIN[1] + (tile_id % columns) * tile_size


def _coords(y0, x0, size):
    y = torch.arange(y0, y0 + size).float(); x = torch.arange(x0, x0 + size).float()
    lat = 90 - (y + .5) * 180 / 720; lon = -180 + (x + .5) * 360 / 1440
    return torch.meshgrid(torch.deg2rad(lat), torch.deg2rad(lon), indexing="ij")


def procedural_state(y0, x0, size, sample, lead=0):
    lat, lon = _coords(y0, x0, size); channel = torch.arange(70).float()[:, None, None]
    phase = .07 * sample + .11 * lead + channel * .031
    return torch.sin((1 + channel % 4) * lat + phase) * torch.cos((1 + channel % 5) * lon - phase) + .3 * torch.sin(3 * lon + 2 * lat + phase)


def make_sample(tile_id, sample, size=32, missing_probability=.2):
    y0, x0 = tile_origin(tile_id, size); target = procedural_state(y0, x0, size, sample)
    lat, lon = _coords(y0, x0, size); channels = torch.arange(70).float()[:, None, None]
    error = .12 * torch.sin(2 * lon + channels * .09) + .05 * torch.cos(3 * lat)
    background = target + error
    times = torch.arange(8).float()[:, None, None, None]; variables = torch.arange(15).long()
    obs = target[(variables * 4) % 70][None] + .35 * background[(variables * 3) % 70][None] + .02 * times
    generator = torch.Generator().manual_seed(100003 * sample + tile_id)
    missing = torch.rand((8, 1, size, size), generator=generator) < missing_probability
    obs = obs.masked_fill(missing.expand_as(obs), 0).reshape(120, size, size)
    forecasts = torch.stack([procedural_state(y0, x0, size, sample, i + 1) for i in range(10)])
    return {"background": background, "obs": obs, "target": target, "correction": background - error, "forecast_targets": forecasts,
            "latitude": torch.rad2deg(lat[:, 0]), "origin": torch.tensor([y0, x0]), "tile_id": torch.tensor(tile_id)}


class ProceduralTileDataset(Dataset):
    def __init__(self, tile_ids, samples_per_tile=2, tile_size=32, missing_probability=.2):
        self.tile_ids, self.samples_per_tile, self.tile_size, self.missing_probability = tuple(tile_ids), samples_per_tile, tile_size, missing_probability
    def __len__(self): return len(self.tile_ids) * self.samples_per_tile
    def __getitem__(self, index):
        return make_sample(self.tile_ids[index % len(self.tile_ids)], index // len(self.tile_ids), self.tile_size, self.missing_probability)


class ChannelLayerNorm(nn.Module):
    def __init__(self, channels):
        super().__init__()
        self.norm = nn.LayerNorm(channels)

    def forward(self, x):
        return self.norm(x.permute(0, 2, 3, 1)).permute(0, 3, 1, 2)


class Down(nn.Sequential):
    def __init__(self, input_channels, output_channels):
        super().__init__(
            nn.Conv2d(input_channels, output_channels, 2, stride=2),
            ChannelLayerNorm(output_channels),
            nn.SiLU(),
            nn.Conv2d(output_channels, output_channels, 3, padding=1),
        )


class Up(nn.Sequential):
    def __init__(self, input_channels, output_channels):
        super().__init__(
            nn.Conv2d(input_channels, input_channels, 3, padding=1),
            ChannelLayerNorm(input_channels),
            nn.SiLU(),
            nn.Conv2d(input_channels, output_channels * 4, 3, padding=1),
            nn.PixelShuffle(2),
        )


class Stem(nn.Sequential):
    def __init__(self, input_channels, channels):
        super().__init__(nn.Conv2d(input_channels, channels, 3, padding=1), ChannelLayerNorm(channels), nn.SiLU())


class Fusion(nn.Module):
    """Interact modalities and emit increment, observation bias, and corrected condition."""
    def __init__(self, channels):
        super().__init__()
        self.mix = nn.Sequential(
            nn.Conv2d(channels * 3, channels, 3, padding=1), ChannelLayerNorm(channels), nn.SiLU(),
            nn.Conv2d(channels, channels * 3, 3, padding=1),
        )

    def forward(self, background, observation, condition):
        increment, obs_bias, corrected_condition = self.mix(torch.cat((background, observation, condition), dim=1)).chunk(3, dim=1)
        return background + increment, observation + obs_bias, condition + corrected_condition


class FuXiDA(nn.Module):
    def __init__(self, base_channels=8):
        super().__init__()
        c = base_channels
        self.background_stem, self.obs_stem, self.condition_stem = Stem(70, c), Stem(120, c), Stem(190, c)
        self.fuse0 = Fusion(c)
        self.background_down1, self.obs_down1, self.condition_down1 = Down(c, 2*c), Down(c, 2*c), Down(c, 2*c)
        self.fuse1 = Fusion(2*c)
        self.background_down2, self.obs_down2, self.condition_down2 = Down(2*c, 4*c), Down(2*c, 4*c), Down(2*c, 4*c)
        self.fuse2 = Fusion(4*c)
        self.up1 = Up(4*c, 2*c)
        self.decode1 = nn.Sequential(nn.Conv2d(4*c, 2*c, 3, padding=1), ChannelLayerNorm(2*c), nn.SiLU())
        self.up0 = Up(2*c, c)
        self.decode0 = nn.Sequential(nn.Conv2d(2*c, c, 3, padding=1), ChannelLayerNorm(c), nn.SiLU())
        self.increment_head = nn.Conv2d(c, 70, 3, padding=1)

    def forward(self, background, observation):
        condition = torch.cat((background, observation), dim=1)
        b0, o0, c0 = self.fuse0(self.background_stem(background), self.obs_stem(observation), self.condition_stem(condition))
        b1, o1, c1 = self.fuse1(self.background_down1(b0), self.obs_down1(o0), self.condition_down1(c0))
        b2, _, _ = self.fuse2(self.background_down2(b1), self.obs_down2(o1), self.condition_down2(c1))
        decoded = self.decode1(torch.cat((self.up1(b2), b1), dim=1))
        decoded = self.decode0(torch.cat((self.up0(decoded), b0), dim=1))
        increment = self.increment_head(decoded)
        return background + increment


class CompactForecastProxy(nn.Module):
    """Frozen differentiable forecast surrogate; gradients only update FuXi-DA."""
    def __init__(self):
        super().__init__()
        self.depthwise = nn.Conv2d(70, 70, 3, padding=1, groups=70, bias=False)
        with torch.no_grad():
            self.depthwise.weight.fill_(1.0 / 9.0)
        self.requires_grad_(False)
        self.eval()

    def forward(self, state):
        return 0.985 * self.depthwise(state) + 0.015 * torch.roll(state, shifts=1, dims=-1)

    def train(self, mode=True):
        return super().train(False)