FuXi-DA / model /fuxi_da.py
zhangrenchao's picture
Publish FuXi-DA reproduction
15aff58 verified
Raw History Blame Contribute Delete
6.79 kB
"""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)