Publish FuXi-DA reproduction
Browse files- .gitattributes +1 -34
- conf/config.yaml +33 -0
- config.json +1 -0
- model/fuxi_da.py +144 -0
- scripts/fake_data.py +8 -0
- scripts/inference.py +9 -0
- scripts/result.py +85 -0
- scripts/train.py +102 -0
- weight/.gitkeep +0 -0
.gitattributes
CHANGED
|
@@ -1,35 +1,2 @@
|
|
| 1 |
-
*.7z filter=lfs diff=lfs merge=lfs -text
|
| 2 |
-
*.arrow filter=lfs diff=lfs merge=lfs -text
|
| 3 |
-
*.bin filter=lfs diff=lfs merge=lfs -text
|
| 4 |
-
*.bz2 filter=lfs diff=lfs merge=lfs -text
|
| 5 |
-
*.ckpt filter=lfs diff=lfs merge=lfs -text
|
| 6 |
-
*.ftz filter=lfs diff=lfs merge=lfs -text
|
| 7 |
-
*.gz filter=lfs diff=lfs merge=lfs -text
|
| 8 |
-
*.h5 filter=lfs diff=lfs merge=lfs -text
|
| 9 |
-
*.joblib filter=lfs diff=lfs merge=lfs -text
|
| 10 |
-
*.lfs.* filter=lfs diff=lfs merge=lfs -text
|
| 11 |
-
*.mlmodel filter=lfs diff=lfs merge=lfs -text
|
| 12 |
-
*.model filter=lfs diff=lfs merge=lfs -text
|
| 13 |
-
*.msgpack filter=lfs diff=lfs merge=lfs -text
|
| 14 |
-
*.npy filter=lfs diff=lfs merge=lfs -text
|
| 15 |
-
*.npz filter=lfs diff=lfs merge=lfs -text
|
| 16 |
-
*.onnx filter=lfs diff=lfs merge=lfs -text
|
| 17 |
-
*.ot filter=lfs diff=lfs merge=lfs -text
|
| 18 |
-
*.parquet filter=lfs diff=lfs merge=lfs -text
|
| 19 |
-
*.pb filter=lfs diff=lfs merge=lfs -text
|
| 20 |
-
*.pickle filter=lfs diff=lfs merge=lfs -text
|
| 21 |
-
*.pkl filter=lfs diff=lfs merge=lfs -text
|
| 22 |
*.pt filter=lfs diff=lfs merge=lfs -text
|
| 23 |
-
*.
|
| 24 |
-
*.rar filter=lfs diff=lfs merge=lfs -text
|
| 25 |
-
*.safetensors filter=lfs diff=lfs merge=lfs -text
|
| 26 |
-
saved_model/**/* filter=lfs diff=lfs merge=lfs -text
|
| 27 |
-
*.tar.* filter=lfs diff=lfs merge=lfs -text
|
| 28 |
-
*.tar filter=lfs diff=lfs merge=lfs -text
|
| 29 |
-
*.tflite filter=lfs diff=lfs merge=lfs -text
|
| 30 |
-
*.tgz filter=lfs diff=lfs merge=lfs -text
|
| 31 |
-
*.wasm filter=lfs diff=lfs merge=lfs -text
|
| 32 |
-
*.xz filter=lfs diff=lfs merge=lfs -text
|
| 33 |
-
*.zip filter=lfs diff=lfs merge=lfs -text
|
| 34 |
-
*.zst filter=lfs diff=lfs merge=lfs -text
|
| 35 |
-
*tfevents* filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
*.pt filter=lfs diff=lfs merge=lfs -text
|
| 2 |
+
*.npz filter=lfs diff=lfs merge=lfs -text
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
conf/config.yaml
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
seed: 2025
|
| 2 |
+
data:
|
| 3 |
+
format_version: fuxi_da_tiles_v1
|
| 4 |
+
manifest: data/fuxi_da_manifest.json
|
| 5 |
+
tile_size: 32
|
| 6 |
+
tile_ids: [0, 37, 181, 399]
|
| 7 |
+
missing_probability: 0.2
|
| 8 |
+
model:
|
| 9 |
+
base_channels: 8
|
| 10 |
+
train:
|
| 11 |
+
iterations: 2
|
| 12 |
+
forecast_steps: 2
|
| 13 |
+
batch_size: 1
|
| 14 |
+
learning_rate: 0.002
|
| 15 |
+
warmup_steps: 1
|
| 16 |
+
weight_decay: 0.00001
|
| 17 |
+
runtime:
|
| 18 |
+
device: cpu
|
| 19 |
+
paths:
|
| 20 |
+
checkpoint: result/checkpoints/fuxi_da.pt
|
| 21 |
+
training_metrics: result/training/metrics.json
|
| 22 |
+
predictions: result/output/predictions.npz
|
| 23 |
+
evaluation: result/evaluation/metrics.json
|
| 24 |
+
figure: result/evaluation/comparison.png
|
| 25 |
+
paper_model:
|
| 26 |
+
background_shape: [70, 721, 1440]
|
| 27 |
+
observation_shape: [8, 15, 640, 640]
|
| 28 |
+
iterations: 6000
|
| 29 |
+
forecast_steps: 10
|
| 30 |
+
batch_size_per_gpu: 1
|
| 31 |
+
gpus: 4
|
| 32 |
+
warmup_steps: 500
|
| 33 |
+
peak_learning_rate: 0.002
|
config.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{"model_name":"FuXi-DA","model_type":"fuxi_da","architectures":["FuXiDA"],"framework":"PyTorch","domain":"weather","task":"satellite-data-assimilation","implementation":{"entry_point":"model/fuxi_da.py","scope":"core-method full-variable tile-streamed engineering reproduction","train_script":"scripts/train.py","inference_script":"scripts/inference.py","evaluation_script":"scripts/result.py","synthetic_data_script":"scripts/fake_data.py"},"architecture":{"background_shape":[70,721,1440],"observation_shape":[8,15,640,640],"analysis_shape":[70,720,1440],"engineering_tile":[32,32],"core":"multi-branch U-Net and multi-scale fusion"},"data":{"datasets":["ERA5","Fengyun-4B AGRI"],"format_version":"fuxi_da_tiles_v1","variables":70,"observation_channels":15,"observation_frames":8,"synthetic":true},"configuration_sources":["conf/config.yaml","model/fuxi_da.py","scripts/fake_data.py","scripts/train.py","scripts/inference.py","scripts/result.py"]}
|
model/fuxi_da.py
ADDED
|
@@ -0,0 +1,144 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Independent PyTorch implementation of a compact multi-branch FuXi-DA U-Net."""
|
| 2 |
+
import torch
|
| 3 |
+
from torch import nn
|
| 4 |
+
import torch.nn.functional as F
|
| 5 |
+
from torch.utils.data import Dataset
|
| 6 |
+
|
| 7 |
+
BACKGROUND_RAW_SHAPE = (70, 721, 1440)
|
| 8 |
+
BACKGROUND_MODEL_SHAPE = (70, 720, 1440)
|
| 9 |
+
AGRI_RAW_SHAPE = (8, 15, 640, 640)
|
| 10 |
+
ANALYSIS_SHAPE = (70, 720, 1440)
|
| 11 |
+
FOOTPRINT_ORIGIN = (40, 400)
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
def tile_origin(tile_id, tile_size=32):
|
| 15 |
+
columns = 640 // tile_size
|
| 16 |
+
return FOOTPRINT_ORIGIN[0] + (tile_id // columns) * tile_size, FOOTPRINT_ORIGIN[1] + (tile_id % columns) * tile_size
|
| 17 |
+
|
| 18 |
+
|
| 19 |
+
def _coords(y0, x0, size):
|
| 20 |
+
y = torch.arange(y0, y0 + size).float(); x = torch.arange(x0, x0 + size).float()
|
| 21 |
+
lat = 90 - (y + .5) * 180 / 720; lon = -180 + (x + .5) * 360 / 1440
|
| 22 |
+
return torch.meshgrid(torch.deg2rad(lat), torch.deg2rad(lon), indexing="ij")
|
| 23 |
+
|
| 24 |
+
|
| 25 |
+
def procedural_state(y0, x0, size, sample, lead=0):
|
| 26 |
+
lat, lon = _coords(y0, x0, size); channel = torch.arange(70).float()[:, None, None]
|
| 27 |
+
phase = .07 * sample + .11 * lead + channel * .031
|
| 28 |
+
return torch.sin((1 + channel % 4) * lat + phase) * torch.cos((1 + channel % 5) * lon - phase) + .3 * torch.sin(3 * lon + 2 * lat + phase)
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
def make_sample(tile_id, sample, size=32, missing_probability=.2):
|
| 32 |
+
y0, x0 = tile_origin(tile_id, size); target = procedural_state(y0, x0, size, sample)
|
| 33 |
+
lat, lon = _coords(y0, x0, size); channels = torch.arange(70).float()[:, None, None]
|
| 34 |
+
error = .12 * torch.sin(2 * lon + channels * .09) + .05 * torch.cos(3 * lat)
|
| 35 |
+
background = target + error
|
| 36 |
+
times = torch.arange(8).float()[:, None, None, None]; variables = torch.arange(15).long()
|
| 37 |
+
obs = target[(variables * 4) % 70][None] + .35 * background[(variables * 3) % 70][None] + .02 * times
|
| 38 |
+
generator = torch.Generator().manual_seed(100003 * sample + tile_id)
|
| 39 |
+
missing = torch.rand((8, 1, size, size), generator=generator) < missing_probability
|
| 40 |
+
obs = obs.masked_fill(missing.expand_as(obs), 0).reshape(120, size, size)
|
| 41 |
+
forecasts = torch.stack([procedural_state(y0, x0, size, sample, i + 1) for i in range(10)])
|
| 42 |
+
return {"background": background, "obs": obs, "target": target, "correction": background - error, "forecast_targets": forecasts,
|
| 43 |
+
"latitude": torch.rad2deg(lat[:, 0]), "origin": torch.tensor([y0, x0]), "tile_id": torch.tensor(tile_id)}
|
| 44 |
+
|
| 45 |
+
|
| 46 |
+
class ProceduralTileDataset(Dataset):
|
| 47 |
+
def __init__(self, tile_ids, samples_per_tile=2, tile_size=32, missing_probability=.2):
|
| 48 |
+
self.tile_ids, self.samples_per_tile, self.tile_size, self.missing_probability = tuple(tile_ids), samples_per_tile, tile_size, missing_probability
|
| 49 |
+
def __len__(self): return len(self.tile_ids) * self.samples_per_tile
|
| 50 |
+
def __getitem__(self, index):
|
| 51 |
+
return make_sample(self.tile_ids[index % len(self.tile_ids)], index // len(self.tile_ids), self.tile_size, self.missing_probability)
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
class ChannelLayerNorm(nn.Module):
|
| 55 |
+
def __init__(self, channels):
|
| 56 |
+
super().__init__()
|
| 57 |
+
self.norm = nn.LayerNorm(channels)
|
| 58 |
+
|
| 59 |
+
def forward(self, x):
|
| 60 |
+
return self.norm(x.permute(0, 2, 3, 1)).permute(0, 3, 1, 2)
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
class Down(nn.Sequential):
|
| 64 |
+
def __init__(self, input_channels, output_channels):
|
| 65 |
+
super().__init__(
|
| 66 |
+
nn.Conv2d(input_channels, output_channels, 2, stride=2),
|
| 67 |
+
ChannelLayerNorm(output_channels),
|
| 68 |
+
nn.SiLU(),
|
| 69 |
+
nn.Conv2d(output_channels, output_channels, 3, padding=1),
|
| 70 |
+
)
|
| 71 |
+
|
| 72 |
+
|
| 73 |
+
class Up(nn.Sequential):
|
| 74 |
+
def __init__(self, input_channels, output_channels):
|
| 75 |
+
super().__init__(
|
| 76 |
+
nn.Conv2d(input_channels, input_channels, 3, padding=1),
|
| 77 |
+
ChannelLayerNorm(input_channels),
|
| 78 |
+
nn.SiLU(),
|
| 79 |
+
nn.Conv2d(input_channels, output_channels * 4, 3, padding=1),
|
| 80 |
+
nn.PixelShuffle(2),
|
| 81 |
+
)
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
class Stem(nn.Sequential):
|
| 85 |
+
def __init__(self, input_channels, channels):
|
| 86 |
+
super().__init__(nn.Conv2d(input_channels, channels, 3, padding=1), ChannelLayerNorm(channels), nn.SiLU())
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
class Fusion(nn.Module):
|
| 90 |
+
"""Interact modalities and emit increment, observation bias, and corrected condition."""
|
| 91 |
+
def __init__(self, channels):
|
| 92 |
+
super().__init__()
|
| 93 |
+
self.mix = nn.Sequential(
|
| 94 |
+
nn.Conv2d(channels * 3, channels, 3, padding=1), ChannelLayerNorm(channels), nn.SiLU(),
|
| 95 |
+
nn.Conv2d(channels, channels * 3, 3, padding=1),
|
| 96 |
+
)
|
| 97 |
+
|
| 98 |
+
def forward(self, background, observation, condition):
|
| 99 |
+
increment, obs_bias, corrected_condition = self.mix(torch.cat((background, observation, condition), dim=1)).chunk(3, dim=1)
|
| 100 |
+
return background + increment, observation + obs_bias, condition + corrected_condition
|
| 101 |
+
|
| 102 |
+
|
| 103 |
+
class FuXiDA(nn.Module):
|
| 104 |
+
def __init__(self, base_channels=8):
|
| 105 |
+
super().__init__()
|
| 106 |
+
c = base_channels
|
| 107 |
+
self.background_stem, self.obs_stem, self.condition_stem = Stem(70, c), Stem(120, c), Stem(190, c)
|
| 108 |
+
self.fuse0 = Fusion(c)
|
| 109 |
+
self.background_down1, self.obs_down1, self.condition_down1 = Down(c, 2*c), Down(c, 2*c), Down(c, 2*c)
|
| 110 |
+
self.fuse1 = Fusion(2*c)
|
| 111 |
+
self.background_down2, self.obs_down2, self.condition_down2 = Down(2*c, 4*c), Down(2*c, 4*c), Down(2*c, 4*c)
|
| 112 |
+
self.fuse2 = Fusion(4*c)
|
| 113 |
+
self.up1 = Up(4*c, 2*c)
|
| 114 |
+
self.decode1 = nn.Sequential(nn.Conv2d(4*c, 2*c, 3, padding=1), ChannelLayerNorm(2*c), nn.SiLU())
|
| 115 |
+
self.up0 = Up(2*c, c)
|
| 116 |
+
self.decode0 = nn.Sequential(nn.Conv2d(2*c, c, 3, padding=1), ChannelLayerNorm(c), nn.SiLU())
|
| 117 |
+
self.increment_head = nn.Conv2d(c, 70, 3, padding=1)
|
| 118 |
+
|
| 119 |
+
def forward(self, background, observation):
|
| 120 |
+
condition = torch.cat((background, observation), dim=1)
|
| 121 |
+
b0, o0, c0 = self.fuse0(self.background_stem(background), self.obs_stem(observation), self.condition_stem(condition))
|
| 122 |
+
b1, o1, c1 = self.fuse1(self.background_down1(b0), self.obs_down1(o0), self.condition_down1(c0))
|
| 123 |
+
b2, _, _ = self.fuse2(self.background_down2(b1), self.obs_down2(o1), self.condition_down2(c1))
|
| 124 |
+
decoded = self.decode1(torch.cat((self.up1(b2), b1), dim=1))
|
| 125 |
+
decoded = self.decode0(torch.cat((self.up0(decoded), b0), dim=1))
|
| 126 |
+
increment = self.increment_head(decoded)
|
| 127 |
+
return background + increment
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
class CompactForecastProxy(nn.Module):
|
| 131 |
+
"""Frozen differentiable forecast surrogate; gradients only update FuXi-DA."""
|
| 132 |
+
def __init__(self):
|
| 133 |
+
super().__init__()
|
| 134 |
+
self.depthwise = nn.Conv2d(70, 70, 3, padding=1, groups=70, bias=False)
|
| 135 |
+
with torch.no_grad():
|
| 136 |
+
self.depthwise.weight.fill_(1.0 / 9.0)
|
| 137 |
+
self.requires_grad_(False)
|
| 138 |
+
self.eval()
|
| 139 |
+
|
| 140 |
+
def forward(self, state):
|
| 141 |
+
return 0.985 * self.depthwise(state) + 0.015 * torch.roll(state, shifts=1, dims=-1)
|
| 142 |
+
|
| 143 |
+
def train(self, mode=True):
|
| 144 |
+
return super().train(False)
|
scripts/fake_data.py
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from pathlib import Path
|
| 2 |
+
import json,sys
|
| 3 |
+
import yaml
|
| 4 |
+
ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
|
| 5 |
+
from model.fuxi_da import *
|
| 6 |
+
c=yaml.safe_load((ROOT/"conf/config.yaml").read_text()); sample=make_sample(c["data"]["tile_ids"][0],0,c["data"]["tile_size"],c["data"]["missing_probability"])
|
| 7 |
+
assert BACKGROUND_RAW_SHAPE==(70,721,1440) and sample["background"].shape==(70,c["data"]["tile_size"],c["data"]["tile_size"]) and sample["obs"].shape[0]==120
|
| 8 |
+
path=ROOT/c["data"]["manifest"];path.parent.mkdir(parents=True,exist_ok=True);path.write_text(json.dumps({"format_version":c["data"]["format_version"],"background_raw_shape":[70,721,1440],"background_model_shape":[70,720,1440],"agri_raw_shape":[8,15,640,640],"observation_model_shape":[120,640,640],"condition_shape":[190,640,640],"analysis_shape":[70,720,1440],"tile_ids":c["data"]["tile_ids"],"tile_size":c["data"]["tile_size"],"is_complete_global":False},indent=2));print(path)
|
scripts/inference.py
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from pathlib import Path
|
| 2 |
+
import sys,yaml,numpy as np,torch
|
| 3 |
+
ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
|
| 4 |
+
from model.fuxi_da import FuXiDA,make_sample
|
| 5 |
+
c=yaml.safe_load((ROOT/"conf/config.yaml").read_text());ck=torch.load(ROOT/c["paths"]["checkpoint"],map_location="cpu",weights_only=True);assert ck["format_version"]==c["data"]["format_version"];m=FuXiDA(**ck["model_config"]);m.load_state_dict(ck["model"]);m.eval();rows=[]
|
| 6 |
+
with torch.no_grad():
|
| 7 |
+
for tile in c["data"]["tile_ids"]:
|
| 8 |
+
s=make_sample(tile,20,c["data"]["tile_size"],c["data"]["missing_probability"]);p=m(s["background"][None],s["obs"][None])[0];rows.append((s,p))
|
| 9 |
+
path=ROOT/c["paths"]["predictions"];path.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(path,background=np.stack([r[0]["background"].numpy() for r in rows]),observations=np.stack([r[0]["obs"].numpy() for r in rows]),target=np.stack([r[0]["target"].numpy() for r in rows]),prediction=np.stack([r[1].numpy() for r in rows]),latitude=np.stack([r[0]["latitude"].numpy() for r in rows]),origins=np.stack([r[0]["origin"].numpy() for r in rows]),tile_ids=np.array(c["data"]["tile_ids"]),format_version=np.array(c["data"]["format_version"]),is_complete_global=np.bool_(False));print(path)
|
scripts/result.py
ADDED
|
@@ -0,0 +1,85 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Tile diagnostics for analysis, forecasts, variable groups, and localization."""
|
| 2 |
+
import argparse
|
| 3 |
+
import json
|
| 4 |
+
from pathlib import Path
|
| 5 |
+
|
| 6 |
+
import matplotlib.pyplot as plt
|
| 7 |
+
import numpy as np
|
| 8 |
+
import torch
|
| 9 |
+
import sys, yaml
|
| 10 |
+
ROOT = Path(__file__).resolve().parents[1]
|
| 11 |
+
sys.path.insert(0, str(ROOT))
|
| 12 |
+
|
| 13 |
+
from model.fuxi_da import CompactForecastProxy, FuXiDA, make_sample
|
| 14 |
+
|
| 15 |
+
GROUPS = {"Z": slice(0, 13), "T": slice(13, 26), "U": slice(26, 39), "V": slice(39, 52), "R": slice(52, 65), "surface": slice(65, 70)}
|
| 16 |
+
|
| 17 |
+
|
| 18 |
+
def weighted_rmse(prediction, target, latitude):
|
| 19 |
+
weight = torch.cos(torch.deg2rad(latitude)).clamp_min(0)
|
| 20 |
+
weight = weight / weight.sum()
|
| 21 |
+
return torch.sqrt(((prediction - target).square() * weight[None, :, None]).sum(dim=(-2, -1)) / prediction.shape[-1]).mean().item()
|
| 22 |
+
|
| 23 |
+
|
| 24 |
+
def main():
|
| 25 |
+
parser = argparse.ArgumentParser()
|
| 26 |
+
parser.add_argument("--config", default="conf/config.yaml")
|
| 27 |
+
parser.add_argument("--checkpoint")
|
| 28 |
+
parser.add_argument("--output")
|
| 29 |
+
parser.add_argument("--metrics")
|
| 30 |
+
args = parser.parse_args()
|
| 31 |
+
cfg = yaml.safe_load((ROOT / args.config).read_text())
|
| 32 |
+
model = FuXiDA(cfg["model"]["base_channels"])
|
| 33 |
+
checkpoint_path = Path(args.checkpoint) if args.checkpoint else ROOT / cfg["paths"]["checkpoint"]
|
| 34 |
+
if checkpoint_path.exists():
|
| 35 |
+
model.load_state_dict(torch.load(checkpoint_path, map_location="cpu", weights_only=True)["model"])
|
| 36 |
+
model.eval()
|
| 37 |
+
proxy = CompactForecastProxy()
|
| 38 |
+
|
| 39 |
+
rows = []
|
| 40 |
+
with torch.no_grad():
|
| 41 |
+
for tile_id in cfg["data"]["tile_ids"]:
|
| 42 |
+
sample = make_sample(tile_id, 20, cfg["data"]["tile_size"], cfg["data"]["missing_probability"])
|
| 43 |
+
analysis = model(sample["background"][None], sample["obs"][None])[0]
|
| 44 |
+
row = {"tile_id": tile_id, "origin": sample["origin"].tolist(), "background": weighted_rmse(sample["background"], sample["target"], sample["latitude"]), "correction": weighted_rmse(sample["correction"], sample["target"], sample["latitude"]), "analysis": weighted_rmse(analysis, sample["target"], sample["latitude"])}
|
| 45 |
+
row["variable_groups"] = {name: weighted_rmse(analysis[index], sample["target"][index], sample["latitude"]) for name, index in GROUPS.items()}
|
| 46 |
+
state = analysis
|
| 47 |
+
row["forecast_steps"] = []
|
| 48 |
+
for lead in range(cfg["train"]["forecast_steps"]):
|
| 49 |
+
state = proxy(state[None])[0]
|
| 50 |
+
row["forecast_steps"].append(weighted_rmse(state, sample["forecast_targets"][lead], sample["latitude"]))
|
| 51 |
+
rows.append(row)
|
| 52 |
+
|
| 53 |
+
sample = make_sample(cfg["data"]["tile_ids"][0], 21, cfg["data"]["tile_size"], 0.0)
|
| 54 |
+
base = model(sample["background"][None], sample["obs"][None])[0]
|
| 55 |
+
perturbed_obs = sample["obs"].clone()
|
| 56 |
+
center = cfg["data"]["tile_size"] // 2
|
| 57 |
+
perturbed_obs[9 - 8, center, center] += 1.0
|
| 58 |
+
response = (model(sample["background"][None], perturbed_obs[None])[0] - base).square().sum(0)
|
| 59 |
+
yy, xx = torch.meshgrid(torch.arange(cfg["data"]["tile_size"]), torch.arange(cfg["data"]["tile_size"]), indexing="ij")
|
| 60 |
+
radius = torch.sqrt((yy - center).float().square() + (xx - center).float().square())
|
| 61 |
+
total_energy = response.sum().clamp_min(1e-12)
|
| 62 |
+
localization = {"perturbation": "AGRI channel 9 +1 K at tile center", "energy_weighted_radius_gridpoints": (response * radius).sum().div(total_energy).item(), "energy_within_radius_4": response[radius <= 4].sum().div(total_energy).item()}
|
| 63 |
+
|
| 64 |
+
metrics = {"coverage_complete": False, "claim": "diagnostic metrics on selected aligned tiles; not complete global scores", "tiles": rows, "increment_localization": localization}
|
| 65 |
+
metrics_path = Path(args.metrics) if args.metrics else ROOT / cfg["paths"]["evaluation"]
|
| 66 |
+
metrics_path.parent.mkdir(parents=True, exist_ok=True); metrics_path.write_text(json.dumps(metrics, indent=2) + "\n")
|
| 67 |
+
labels = ["background", "correction", "analysis"]
|
| 68 |
+
values = [np.mean([row[label] for row in rows]) for label in labels]
|
| 69 |
+
fig, axes = plt.subplots(1, 2, figsize=(10, 4))
|
| 70 |
+
axes[0].bar(labels, values, color=["#596780", "#d69b36", "#287f71"])
|
| 71 |
+
axes[0].set_ylabel("Latitude-weighted RMSE")
|
| 72 |
+
axes[0].set_title("Selected aligned tiles (incomplete coverage)")
|
| 73 |
+
group_values = [np.mean([row["variable_groups"][name] for row in rows]) for name in GROUPS]
|
| 74 |
+
axes[1].bar(list(GROUPS), group_values, color="#287f71")
|
| 75 |
+
axes[1].set_title("Analysis RMSE by variable group")
|
| 76 |
+
fig.suptitle("FuXi-DA procedural protocol diagnostics")
|
| 77 |
+
fig.tight_layout()
|
| 78 |
+
output_path = Path(args.output) if args.output else ROOT / cfg["paths"]["figure"]
|
| 79 |
+
output_path.parent.mkdir(parents=True, exist_ok=True); fig.savefig(output_path, dpi=160)
|
| 80 |
+
plt.close(fig)
|
| 81 |
+
print(json.dumps(metrics, indent=2))
|
| 82 |
+
|
| 83 |
+
|
| 84 |
+
if __name__ == "__main__":
|
| 85 |
+
main()
|
scripts/train.py
ADDED
|
@@ -0,0 +1,102 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""DDP-capable FuXi-DA training with analysis and frozen-proxy forecast supervision."""
|
| 2 |
+
import argparse
|
| 3 |
+
import json
|
| 4 |
+
import math
|
| 5 |
+
import os
|
| 6 |
+
from pathlib import Path
|
| 7 |
+
|
| 8 |
+
import torch
|
| 9 |
+
import torch.distributed as dist
|
| 10 |
+
from torch.nn.parallel import DistributedDataParallel
|
| 11 |
+
from torch.utils.data import DataLoader, DistributedSampler
|
| 12 |
+
|
| 13 |
+
import sys
|
| 14 |
+
ROOT = Path(__file__).resolve().parents[1]
|
| 15 |
+
sys.path.insert(0, str(ROOT))
|
| 16 |
+
import yaml
|
| 17 |
+
from model.fuxi_da import CompactForecastProxy, FuXiDA, ProceduralTileDataset
|
| 18 |
+
|
| 19 |
+
|
| 20 |
+
def latitude_weighted_l1(prediction, target, latitude):
|
| 21 |
+
weight = torch.cos(torch.deg2rad(latitude)).clamp_min(0)
|
| 22 |
+
weight = weight * weight.shape[-1] / weight.sum(dim=-1, keepdim=True)
|
| 23 |
+
return ((prediction - target).abs() * weight[:, None, :, None]).mean()
|
| 24 |
+
|
| 25 |
+
|
| 26 |
+
def learning_rate(step, warmup, total, peak):
|
| 27 |
+
if step < warmup:
|
| 28 |
+
return 1e-8 + (peak - 1e-8) * step / max(1, warmup)
|
| 29 |
+
progress = (step - warmup) / max(1, total - warmup)
|
| 30 |
+
return peak * 0.5 * (1.0 + math.cos(math.pi * progress))
|
| 31 |
+
|
| 32 |
+
|
| 33 |
+
def main():
|
| 34 |
+
parser = argparse.ArgumentParser()
|
| 35 |
+
parser.add_argument("--config", default="conf/config.yaml")
|
| 36 |
+
parser.add_argument("--iterations", type=int)
|
| 37 |
+
parser.add_argument("--forecast-steps", type=int)
|
| 38 |
+
parser.add_argument("--output-dir")
|
| 39 |
+
args = parser.parse_args()
|
| 40 |
+
cfg = yaml.safe_load((ROOT / args.config).read_text())
|
| 41 |
+
total = args.iterations or cfg["train"]["iterations"]
|
| 42 |
+
forecast_steps = args.forecast_steps or cfg["train"]["forecast_steps"]
|
| 43 |
+
|
| 44 |
+
world_size = int(os.environ.get("WORLD_SIZE", "1"))
|
| 45 |
+
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
|
| 46 |
+
use_cuda = torch.cuda.is_available() and torch.cuda.device_count() >= world_size and cfg["runtime"]["device"] != "cpu"
|
| 47 |
+
if world_size > 1:
|
| 48 |
+
dist.init_process_group("nccl" if use_cuda else "gloo")
|
| 49 |
+
device = torch.device(f"cuda:{local_rank}" if use_cuda else "cpu")
|
| 50 |
+
if use_cuda:
|
| 51 |
+
torch.cuda.set_device(device)
|
| 52 |
+
torch.manual_seed(cfg["seed"] + local_rank)
|
| 53 |
+
|
| 54 |
+
dataset = ProceduralTileDataset(cfg["data"]["tile_ids"], max(2, total), cfg["data"]["tile_size"], cfg["data"]["missing_probability"])
|
| 55 |
+
sampler = DistributedSampler(dataset, shuffle=True) if world_size > 1 else None
|
| 56 |
+
loader = DataLoader(dataset, batch_size=cfg["train"]["batch_size"], sampler=sampler, shuffle=sampler is None, num_workers=0)
|
| 57 |
+
model = FuXiDA(cfg["model"]["base_channels"]).to(device)
|
| 58 |
+
proxy = CompactForecastProxy().to(device)
|
| 59 |
+
if world_size > 1:
|
| 60 |
+
model = DistributedDataParallel(model, device_ids=[local_rank] if device.type == "cuda" else None)
|
| 61 |
+
optimizer = torch.optim.AdamW(model.parameters(), lr=cfg["train"]["learning_rate"], betas=(0.9, 0.999), weight_decay=cfg["train"]["weight_decay"])
|
| 62 |
+
|
| 63 |
+
model.train()
|
| 64 |
+
iterator = iter(loader)
|
| 65 |
+
for step in range(total):
|
| 66 |
+
try:
|
| 67 |
+
batch = next(iterator)
|
| 68 |
+
except StopIteration:
|
| 69 |
+
if sampler is not None:
|
| 70 |
+
sampler.set_epoch(step)
|
| 71 |
+
iterator = iter(loader)
|
| 72 |
+
batch = next(iterator)
|
| 73 |
+
background, obs, target = (batch[key].to(device) for key in ("background", "obs", "target"))
|
| 74 |
+
latitude = batch["latitude"].to(device)
|
| 75 |
+
analysis = model(background, obs)
|
| 76 |
+
analysis_loss = latitude_weighted_l1(analysis, target, latitude)
|
| 77 |
+
state, forecast_loss = analysis, analysis_loss.new_zeros(())
|
| 78 |
+
for lead in range(forecast_steps):
|
| 79 |
+
state = proxy(state)
|
| 80 |
+
forecast_loss = forecast_loss + latitude_weighted_l1(state, batch["forecast_targets"][:, lead].to(device), latitude)
|
| 81 |
+
loss = analysis_loss + forecast_loss / forecast_steps
|
| 82 |
+
optimizer.zero_grad(set_to_none=True)
|
| 83 |
+
loss.backward()
|
| 84 |
+
optimizer.step()
|
| 85 |
+
lr = learning_rate(step + 1, cfg["train"]["warmup_steps"], total, cfg["train"]["learning_rate"])
|
| 86 |
+
for group in optimizer.param_groups:
|
| 87 |
+
group["lr"] = lr
|
| 88 |
+
if local_rank == 0 and (step == 0 or (step + 1) % 100 == 0 or step + 1 == total):
|
| 89 |
+
print(json.dumps({"step": step + 1, "loss": loss.item(), "analysis_l1": analysis_loss.item(), "lr": lr}))
|
| 90 |
+
|
| 91 |
+
if local_rank == 0:
|
| 92 |
+
checkpoint = ROOT / cfg["paths"]["checkpoint"]; checkpoint.parent.mkdir(parents=True, exist_ok=True)
|
| 93 |
+
bare_model = model.module if isinstance(model, DistributedDataParallel) else model
|
| 94 |
+
torch.save({"model": bare_model.state_dict(), "model_config": cfg["model"], "format_version": cfg["data"]["format_version"]}, checkpoint)
|
| 95 |
+
metrics = ROOT / cfg["paths"]["training_metrics"]; metrics.parent.mkdir(parents=True, exist_ok=True)
|
| 96 |
+
metrics.write_text(json.dumps({"iterations": total, "final_loss": float(loss), "world_size": world_size}, indent=2))
|
| 97 |
+
if world_size > 1:
|
| 98 |
+
dist.destroy_process_group()
|
| 99 |
+
|
| 100 |
+
|
| 101 |
+
if __name__ == "__main__":
|
| 102 |
+
main()
|
weight/.gitkeep
ADDED
|
File without changes
|