zhangrenchao commited on
Commit
15aff58
·
verified ·
1 Parent(s): a31b7c5

Publish FuXi-DA reproduction

Browse files
.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
- *.pth filter=lfs diff=lfs merge=lfs -text
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