zhangrenchao commited on
Commit
380b161
·
verified ·
1 Parent(s): a31b7c5

Publish ACE2 reproduction

Browse files
.gitattributes CHANGED
@@ -1,35 +1,5 @@
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
+ *.json text eol=lf
2
+ *.md text eol=lf
3
+ *.py text eol=lf
4
+ *.yaml text eol=lf
5
+ weight/** filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
conf/config.yaml ADDED
@@ -0,0 +1,33 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ seed: 17
2
+ data:
3
+ format_version: ace2-full-grid-v1
4
+ path: data/ace2_tiny.npz
5
+ samples: 4
6
+ time_steps: 6
7
+ channels: 50
8
+ height: 180
9
+ width: 360
10
+ levels: 8
11
+ dt_hours: 6
12
+ model:
13
+ width: 4
14
+ depth: 8
15
+ modes_lat: 4
16
+ modes_lon: 4
17
+ forcing_channels: 4
18
+ paper_model:
19
+ embedding_dim: 384
20
+ parameters: 450000000
21
+ autoregressive_training_steps: 2
22
+ train:
23
+ epochs: 1
24
+ batch_size: 1
25
+ learning_rate: 0.001
26
+ two_step_weight: 0.5
27
+ checkpoint: result/checkpoints/ace2.pt
28
+ inference:
29
+ steps: 3
30
+ output: result/output/predictions.npz
31
+ evaluation:
32
+ output: result/evaluation/metrics.json
33
+ figure: result/evaluation/comparison.png
config.json ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "model_name": "ACE2",
3
+ "model_type": "ace2",
4
+ "architectures": ["CompactSFNO"],
5
+ "framework": "PyTorch",
6
+ "domain": "climate",
7
+ "task": "global-atmospheric-emulation",
8
+ "implementation": {"entry_point": "model/ace2.py", "scope": "core-method full-grid reduced-model engineering reproduction", "train_script": "scripts/train.py", "inference_script": "scripts/inference.py", "evaluation_script": "scripts/result.py", "synthetic_data_script": "scripts/fake_data.py"},
9
+ "architecture": {"grid": [180,360], "state_channels": 50, "forcing_channels": 4, "vertical_layers": 8, "engineering_width": 4, "engineering_depth": 8, "paper_embedding": 384, "paper_parameters": 450000000},
10
+ "data": {"datasets": ["ERA5","SHiELD"], "format_version": "ace2-full-grid-v1", "time_step_hours": 6, "layout": "BCHW", "synthetic": true},
11
+ "configuration_sources": ["conf/config.yaml","model/ace2.py","scripts/fake_data.py","scripts/train.py","scripts/inference.py","scripts/result.py"]
12
+ }
model/ace2.py ADDED
@@ -0,0 +1,177 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import random
3
+ from pathlib import Path
4
+
5
+ import numpy as np
6
+ import torch
7
+ import yaml
8
+ from torch import nn
9
+
10
+
11
+ LEVELS_HPA = (1000, 850, 700, 500, 300, 200, 100, 50)
12
+ CHANNELS = (
13
+ [f"temperature_{p}" for p in LEVELS_HPA]
14
+ + [f"specific_humidity_{p}" for p in LEVELS_HPA]
15
+ + [f"u_wind_{p}" for p in LEVELS_HPA]
16
+ + [f"v_wind_{p}" for p in LEVELS_HPA]
17
+ + [f"geopotential_{p}" for p in LEVELS_HPA]
18
+ + [
19
+ "surface_pressure",
20
+ "air_temperature_2m",
21
+ "specific_humidity_2m",
22
+ "eastward_wind_10m",
23
+ "northward_wind_10m",
24
+ "sea_surface_temperature",
25
+ "total_precipitation_6h",
26
+ "surface_downward_shortwave",
27
+ "surface_downward_longwave",
28
+ "toa_outgoing_longwave",
29
+ ]
30
+ )
31
+ assert len(CHANNELS) == 50
32
+ Q_INDICES = tuple(range(8, 16)) + (42,)
33
+ SURFACE_PRESSURE = 40
34
+ PRECIPITATION = 46
35
+ RADIATION_INDICES = (47, 48, 49)
36
+
37
+
38
+ class SpectralConv2d(nn.Module):
39
+ def __init__(self, width, modes_lat, modes_lon):
40
+ super().__init__()
41
+ self.modes_lat, self.modes_lon = modes_lat, modes_lon
42
+ scale = 1.0 / width
43
+ self.weight = nn.Parameter(
44
+ scale * torch.randn(width, width, modes_lat, modes_lon, dtype=torch.cfloat)
45
+ )
46
+
47
+ def forward(self, x):
48
+ spectrum = torch.fft.rfft2(x, norm="ortho")
49
+ out = torch.zeros_like(spectrum)
50
+ ml = min(self.modes_lat, spectrum.shape[-2])
51
+ mn = min(self.modes_lon, spectrum.shape[-1])
52
+ out[:, :, :ml, :mn] = torch.einsum(
53
+ "bixy,ioxy->boxy", spectrum[:, :, :ml, :mn], self.weight[:, :, :ml, :mn]
54
+ )
55
+ return torch.fft.irfft2(out, s=x.shape[-2:], norm="ortho")
56
+
57
+
58
+ class SFNOBlock(nn.Module):
59
+ def __init__(self, width, modes_lat, modes_lon):
60
+ super().__init__()
61
+ self.spectral = SpectralConv2d(width, modes_lat, modes_lon)
62
+ self.mlp = nn.Sequential(
63
+ nn.Conv2d(width, width * 2, 1), nn.GELU(), nn.Conv2d(width * 2, width, 1)
64
+ )
65
+ self.norm = nn.GroupNorm(1, width)
66
+
67
+ def forward(self, x):
68
+ return x + self.mlp(self.norm(self.spectral(x)))
69
+
70
+
71
+ class CompactSFNO(nn.Module):
72
+ def __init__(self, channels=50, forcing_channels=4, width=4, depth=1,
73
+ modes_lat=4, modes_lon=4):
74
+ super().__init__()
75
+ self.lift = nn.Conv2d(channels + forcing_channels, width, 1)
76
+ self.blocks = nn.Sequential(
77
+ *[SFNOBlock(width, modes_lat, modes_lon) for _ in range(depth)]
78
+ )
79
+ self.project = nn.Sequential(nn.GELU(), nn.Conv2d(width, channels, 1))
80
+
81
+ def forward(self, state, forcing):
82
+ features = self.blocks(self.lift(torch.cat((state, forcing), dim=1)))
83
+ return state + self.project(features)
84
+
85
+
86
+ def area_weights(height, device, dtype):
87
+ lat = torch.linspace(-89.5, 89.5, height, device=device, dtype=dtype)
88
+ return torch.cos(torch.deg2rad(lat)).view(1, 1, height, 1)
89
+
90
+
91
+ def weighted_mean(x, weights):
92
+ return (x * weights).sum(dim=(-2, -1), keepdim=True) / (
93
+ weights.sum(dim=(-2, -1), keepdim=True) * x.shape[-1]
94
+ )
95
+
96
+
97
+ def hard_correct(previous, predicted):
98
+ """Apply differentiable positivity, dry-mass, and global-water constraints."""
99
+ out = predicted.clone()
100
+ positive = list(Q_INDICES) + [PRECIPITATION] + list(RADIATION_INDICES)
101
+ out[:, positive] = torch.clamp_min(out[:, positive], 0.0)
102
+ weights = area_weights(out.shape[-2], out.device, out.dtype)
103
+
104
+ q_prev = previous[:, Q_INDICES].sum(dim=1, keepdim=True)
105
+ water_target = weighted_mean(q_prev, weights)
106
+ precip = weighted_mean(out[:, PRECIPITATION:PRECIPITATION + 1], weights)
107
+ precip_scale = torch.clamp(
108
+ 0.5 * water_target / torch.clamp_min(precip, 1e-8), max=1.0
109
+ )
110
+ out[:, PRECIPITATION:PRECIPITATION + 1] *= precip_scale
111
+ precip = weighted_mean(out[:, PRECIPITATION:PRECIPITATION + 1], weights)
112
+ q_target = torch.clamp_min(water_target - precip, 0.0)
113
+ q_now = weighted_mean(out[:, Q_INDICES].sum(dim=1, keepdim=True), weights)
114
+ out[:, Q_INDICES] *= q_target / torch.clamp_min(q_now, 1e-8)
115
+
116
+ q_new = out[:, Q_INDICES].sum(dim=1, keepdim=True)
117
+ dry_target = weighted_mean(
118
+ previous[:, SURFACE_PRESSURE:SURFACE_PRESSURE + 1] - q_prev, weights
119
+ )
120
+ dry_now = weighted_mean(
121
+ out[:, SURFACE_PRESSURE:SURFACE_PRESSURE + 1] - q_new, weights
122
+ )
123
+ out[:, SURFACE_PRESSURE:SURFACE_PRESSURE + 1] += dry_target - dry_now
124
+ return out
125
+
126
+
127
+ def load_config(root=None):
128
+ root = Path(root) if root is not None else Path(__file__).resolve().parents[1]
129
+ with (root / "conf" / "config.yaml").open(encoding="utf-8") as handle:
130
+ return yaml.safe_load(handle)
131
+
132
+
133
+ def seed_all(seed):
134
+ random.seed(seed)
135
+ np.random.seed(seed)
136
+ torch.manual_seed(seed)
137
+
138
+
139
+ def forcing_for_hours(hours, height=180, width=360):
140
+ hours = np.asarray(hours, dtype=np.float32)
141
+ phase = 2 * np.pi * hours / (365.25 * 24)
142
+ lat = np.deg2rad(np.linspace(-89.5, 89.5, height, dtype=np.float32))
143
+ lon = np.deg2rad(np.linspace(0.5, 359.5, width, dtype=np.float32))
144
+ solar = np.maximum(
145
+ 0,
146
+ np.cos(lat)[None, :, None]
147
+ * np.cos(lon[None, None, :] + phase[:, None, None]),
148
+ )
149
+ fields = np.empty((len(hours), 4, height, width), dtype=np.float32)
150
+ fields[:, 0] = np.sin(phase)[:, None, None]
151
+ fields[:, 1] = np.cos(phase)[:, None, None]
152
+ fields[:, 2] = (400.0 + 0.01 * hours)[:, None, None] / 500.0
153
+ fields[:, 3] = solar
154
+ return fields
155
+
156
+
157
+ def init_distributed():
158
+ distributed = int(os.environ.get("WORLD_SIZE", "1")) > 1
159
+ world_size = int(os.environ.get("WORLD_SIZE", "1"))
160
+ use_cuda = torch.cuda.is_available() and torch.cuda.device_count() >= world_size
161
+ if distributed:
162
+ backend = "nccl" if use_cuda else "gloo"
163
+ torch.distributed.init_process_group(backend=backend)
164
+ rank = torch.distributed.get_rank()
165
+ local_rank = int(os.environ.get("LOCAL_RANK", "0"))
166
+ else:
167
+ rank = local_rank = 0
168
+ device = torch.device(
169
+ f"cuda:{local_rank}" if use_cuda else "cpu"
170
+ )
171
+ if device.type == "cuda":
172
+ torch.cuda.set_device(device)
173
+ return distributed, rank, device
174
+
175
+
176
+ def build_model(config):
177
+ return CompactSFNO(channels=config["data"]["channels"], **config["model"])
scripts/fake_data.py ADDED
@@ -0,0 +1,50 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ import sys
3
+
4
+ import numpy as np
5
+
6
+ sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
7
+ from model.ace2 import (
8
+ PRECIPITATION,
9
+ Q_INDICES,
10
+ RADIATION_INDICES,
11
+ SURFACE_PRESSURE,
12
+ forcing_for_hours,
13
+ load_config,
14
+ seed_all,
15
+ )
16
+
17
+
18
+ ROOT = Path(__file__).resolve().parents[1]
19
+
20
+
21
+ def main():
22
+ cfg = load_config(ROOT)
23
+ seed_all(cfg["seed"])
24
+ d = cfg["data"]
25
+ shape = (d["samples"], d["time_steps"], d["channels"], d["height"], d["width"])
26
+ rng = np.random.default_rng(cfg["seed"])
27
+ lat = np.deg2rad(np.linspace(-89.5, 89.5, d["height"], dtype=np.float32))[None, None, None, :, None]
28
+ lon = np.deg2rad(np.linspace(0.5, 359.5, d["width"], dtype=np.float32))[None, None, None, None, :]
29
+ time = np.arange(d["time_steps"], dtype=np.float32)[None, :, None, None, None]
30
+ channel = np.arange(d["channels"], dtype=np.float32)[None, None, :, None, None]
31
+ state = (0.2 * np.sin(lat * (1 + channel % 3)) + 0.1 * np.cos(lon + time / 4)
32
+ + 0.002 * channel + 0.003 * time).astype(np.float32)
33
+ state = np.broadcast_to(state, shape).copy()
34
+ state += rng.normal(0, 0.005, (d["samples"], 1, d["channels"], 1, 1)).astype(np.float32)
35
+ state[:, :, Q_INDICES] = np.maximum(state[:, :, Q_INDICES] * 0.01 + 0.003, 0)
36
+ state[:, :, SURFACE_PRESSURE] = 1.0 + 0.01 * np.cos(lat[:, :, 0])
37
+ state[:, :, PRECIPITATION] = np.maximum(0, 0.001 * (np.sin(lon[:, :, 0] + time[:, :, 0]) + 1))
38
+ state[:, :, RADIATION_INDICES] = np.maximum(state[:, :, RADIATION_INDICES] + 0.5, 0)
39
+ hours = (np.arange(d["samples"])[:, None] * d["time_steps"] + np.arange(d["time_steps"])[None]) * d["dt_hours"]
40
+ forcing = forcing_for_hours(hours.reshape(-1), d["height"], d["width"]).reshape(
41
+ d["samples"], d["time_steps"], 4, d["height"], d["width"]
42
+ )
43
+ path = ROOT / d["path"]
44
+ path.parent.mkdir(parents=True, exist_ok=True)
45
+ np.savez_compressed(path, state=state.astype(np.float16), forcing=forcing.astype(np.float16), hours=hours)
46
+ print(f"created {path}: state={state.shape}, forcing={forcing.shape}, dt={d['dt_hours']}h")
47
+
48
+
49
+ if __name__ == "__main__":
50
+ main()
scripts/inference.py ADDED
@@ -0,0 +1,38 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+
7
+ ROOT = Path(__file__).resolve().parents[1]
8
+ sys.path.insert(0, str(ROOT))
9
+ from model.ace2 import build_model, hard_correct, load_config
10
+
11
+
12
+ def main():
13
+ cfg = load_config(ROOT)
14
+ data = np.load(ROOT / cfg["data"]["path"])
15
+ model = build_model(cfg)
16
+ checkpoint = torch.load(
17
+ ROOT / cfg["train"]["checkpoint"], map_location="cpu", weights_only=True
18
+ )
19
+ if checkpoint["format_version"] != cfg["data"]["format_version"]:
20
+ raise ValueError("checkpoint format_version mismatch")
21
+ model.load_state_dict(checkpoint["model"])
22
+ model.eval()
23
+ current = torch.from_numpy(data["state"][0, 0].astype(np.float32)).unsqueeze(0)
24
+ forecast = []
25
+ with torch.no_grad():
26
+ for step in range(cfg["inference"]["steps"]):
27
+ forcing = torch.from_numpy(data["forcing"][0, step + 1].astype(np.float32)).unsqueeze(0)
28
+ current = hard_correct(current, model(current, forcing))
29
+ forecast.append(current.squeeze(0).numpy().astype(np.float16))
30
+ output = ROOT / cfg["inference"]["output"]
31
+ output.parent.mkdir(parents=True, exist_ok=True)
32
+ leads = np.arange(1, cfg["inference"]["steps"] + 1) * cfg["data"]["dt_hours"]
33
+ np.savez_compressed(output, forecast=np.stack(forecast), lead_hours=leads)
34
+ print(f"saved {output}: forecast={tuple(np.stack(forecast).shape)}, leads={leads.tolist()}h")
35
+
36
+
37
+ if __name__ == "__main__":
38
+ main()
scripts/result.py ADDED
@@ -0,0 +1,68 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import json
2
+ from pathlib import Path
3
+ import sys
4
+
5
+ import numpy as np
6
+ import matplotlib
7
+ matplotlib.use("Agg")
8
+ import matplotlib.pyplot as plt
9
+
10
+ ROOT = Path(__file__).resolve().parents[1]
11
+ sys.path.insert(0, str(ROOT))
12
+ from model.ace2 import PRECIPITATION, Q_INDICES, SURFACE_PRESSURE, load_config
13
+
14
+
15
+ def weighted_mean(x, w):
16
+ return np.sum(x * w, axis=(-2, -1)) / np.sum(np.broadcast_to(w, x.shape), axis=(-2, -1))
17
+
18
+
19
+ def rmse(pred, target, w):
20
+ return float(np.sqrt(np.sum((pred - target) ** 2 * w) / np.sum(np.broadcast_to(w, pred.shape))))
21
+
22
+
23
+ def main():
24
+ cfg = load_config(ROOT)
25
+ data = np.load(ROOT / cfg["data"]["path"])
26
+ steps = cfg["inference"]["steps"]
27
+ truth = data["state"][0, 1:steps + 1].astype(np.float32)
28
+ pred = np.load(ROOT / cfg["inference"]["output"])["forecast"].astype(np.float32)
29
+ initial = data["state"][0, 0].astype(np.float32)
30
+ persistence = np.broadcast_to(initial, truth.shape)
31
+ w = np.cos(np.deg2rad(np.linspace(-89.5, 89.5, 180, dtype=np.float32)))[None, None, :, None]
32
+ pg, tg = weighted_mean(pred, w), weighted_mean(truth, w)
33
+ denom = np.sum((tg - tg.mean(axis=0, keepdims=True)) ** 2)
34
+ r2 = float(1 - np.sum((pg - tg) ** 2) / max(float(denom), 1e-12))
35
+ dry0 = weighted_mean(initial[SURFACE_PRESSURE] - initial[list(Q_INDICES)].sum(0), w[0, 0])
36
+ dry = weighted_mean(pred[:, SURFACE_PRESSURE] - pred[:, list(Q_INDICES)].sum(1), w[0, 0])
37
+ previous = np.concatenate((initial[None], pred[:-1]), axis=0)
38
+ water_previous = weighted_mean(previous[:, list(Q_INDICES)].sum(1), w[0, 0])
39
+ water = weighted_mean(pred[:, list(Q_INDICES)].sum(1) + pred[:, PRECIPITATION], w[0, 0])
40
+ model_rmse, baseline_rmse = rmse(pred, truth, w), rmse(persistence, truth, w)
41
+ metrics = {
42
+ "area_weighted_rmse": model_rmse,
43
+ "global_mean_r2": r2,
44
+ "conservation": {
45
+ "max_abs_global_dry_mass_error": float(np.max(np.abs(dry - dry0))),
46
+ "max_abs_global_moisture_closure_error": float(np.max(np.abs(water - water_previous))),
47
+ },
48
+ "comparison": {
49
+ "persistence_area_weighted_rmse": baseline_rmse,
50
+ "rmse_skill_vs_persistence": float(1 - model_rmse / baseline_rmse),
51
+ },
52
+ }
53
+ output = ROOT / cfg["evaluation"]["output"]
54
+ output.parent.mkdir(parents=True, exist_ok=True)
55
+ output.write_text(json.dumps(metrics, indent=2) + "\n", encoding="utf-8")
56
+ figure = ROOT / cfg["evaluation"]["figure"]
57
+ figure.parent.mkdir(parents=True, exist_ok=True)
58
+ fig, axes = plt.subplots(1, 2, figsize=(9, 3.8), constrained_layout=True)
59
+ axes[0].bar(["ACE2", "Persistence"], [model_rmse, baseline_rmse], color=["#287f71", "#8c96a8"])
60
+ axes[0].set(ylabel="Area-weighted RMSE", title="Forecast error")
61
+ axes[1].bar(["Dry mass", "Moisture"], [metrics["conservation"]["max_abs_global_dry_mass_error"], metrics["conservation"]["max_abs_global_moisture_closure_error"]], color="#d69b36")
62
+ axes[1].set_yscale("log"); axes[1].set(title="Conservation residual", ylabel="Maximum absolute error")
63
+ fig.savefig(figure, dpi=150); plt.close(fig)
64
+ print(json.dumps(metrics, indent=2))
65
+
66
+
67
+ if __name__ == "__main__":
68
+ main()
scripts/train.py ADDED
@@ -0,0 +1,72 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ import sys
3
+
4
+ import numpy as np
5
+ import torch
6
+ from torch.nn.parallel import DistributedDataParallel
7
+ from torch.utils.data import DataLoader, Dataset, DistributedSampler
8
+
9
+ ROOT = Path(__file__).resolve().parents[1]
10
+ sys.path.insert(0, str(ROOT))
11
+ from model.ace2 import build_model, hard_correct, init_distributed, load_config, seed_all
12
+
13
+
14
+ class Windows(Dataset):
15
+ def __init__(self, path):
16
+ data = np.load(path)
17
+ self.state = data["state"]
18
+ self.forcing = data["forcing"]
19
+ self.windows = [(n, t) for n in range(len(self.state)) for t in range(self.state.shape[1] - 2)]
20
+
21
+ def __len__(self):
22
+ return len(self.windows)
23
+
24
+ def __getitem__(self, index):
25
+ n, t = self.windows[index]
26
+ return tuple(torch.from_numpy(x.astype(np.float32)) for x in (
27
+ self.state[n, t], self.state[n, t + 1], self.state[n, t + 2],
28
+ self.forcing[n, t + 1], self.forcing[n, t + 2]))
29
+
30
+
31
+ def main():
32
+ cfg = load_config(ROOT)
33
+ seed_all(cfg["seed"])
34
+ distributed, rank, device = init_distributed()
35
+ dataset = Windows(ROOT / cfg["data"]["path"])
36
+ sampler = DistributedSampler(dataset, shuffle=True) if distributed else None
37
+ loader = DataLoader(dataset, batch_size=cfg["train"]["batch_size"], sampler=sampler,
38
+ shuffle=sampler is None, num_workers=0)
39
+ model = build_model(cfg).to(device)
40
+ if distributed:
41
+ model = DistributedDataParallel(model, device_ids=[device.index] if device.type == "cuda" else None)
42
+ optimizer = torch.optim.Adam(model.parameters(), lr=cfg["train"]["learning_rate"])
43
+ for epoch in range(cfg["train"]["epochs"]):
44
+ if sampler:
45
+ sampler.set_epoch(epoch)
46
+ total = 0.0
47
+ for x0, y1, y2, f1, f2 in loader:
48
+ x0, y1, y2, f1, f2 = (x.to(device) for x in (x0, y1, y2, f1, f2))
49
+ p1 = hard_correct(x0, model(x0, f1))
50
+ p2 = hard_correct(p1, model(p1, f2))
51
+ loss = torch.mean((p1 - y1) ** 2) + cfg["train"]["two_step_weight"] * torch.mean((p2 - y2) ** 2)
52
+ optimizer.zero_grad(set_to_none=True)
53
+ loss.backward()
54
+ optimizer.step()
55
+ total += loss.item()
56
+ if rank == 0:
57
+ print(f"epoch={epoch + 1} two_step_loss={total / len(loader):.7f}")
58
+ if rank == 0:
59
+ path = ROOT / cfg["train"]["checkpoint"]
60
+ path.parent.mkdir(parents=True, exist_ok=True)
61
+ raw_model = model.module if distributed else model
62
+ torch.save({"model": raw_model.state_dict(), "model_config": cfg["model"], "format_version": cfg["data"]["format_version"]}, path)
63
+ metrics = ROOT / "result/training/metrics.json"
64
+ metrics.parent.mkdir(parents=True, exist_ok=True)
65
+ metrics.write_text(__import__("json").dumps({"history": [{"epoch": cfg["train"]["epochs"], "loss": total / len(loader)}], "world_size": int(__import__("os").environ.get("WORLD_SIZE", "1"))}, indent=2))
66
+ print(f"saved {path}")
67
+ if distributed:
68
+ torch.distributed.destroy_process_group()
69
+
70
+
71
+ if __name__ == "__main__":
72
+ main()
weight/.gitkeep ADDED
File without changes