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)
|