Download model/fuxi_da.py from OneScience-Group/FuXi-DA: direct link, hf CLI and curl.
- Browser
- Download file 6.79 kB
-
https://huggingface.co/OneScience-Group/FuXi-DA/resolve/main/model/fuxi_da.py
- Command line
-
hf download hf://OneScience-Group/FuXi-DA/model/fuxi_da.py
-
curl -L -o fuxi_da.py https://huggingface.co/OneScience-Group/FuXi-DA/resolve/main/model/fuxi_da.py
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) | |