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