Publish ACE2 reproduction
Browse files- .gitattributes +5 -35
- conf/config.yaml +33 -0
- config.json +12 -0
- model/ace2.py +177 -0
- scripts/fake_data.py +50 -0
- scripts/inference.py +38 -0
- scripts/result.py +68 -0
- scripts/train.py +72 -0
- weight/.gitkeep +0 -0
.gitattributes
CHANGED
|
@@ -1,35 +1,5 @@
|
|
| 1 |
-
*.
|
| 2 |
-
*.
|
| 3 |
-
*.
|
| 4 |
-
*.
|
| 5 |
-
*
|
| 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
|