Publish Climate2Weather reproduction
Browse files- .gitattributes +2 -35
- conf/config.yaml +7 -0
- config.json +1 -0
- model/climate2weather.py +15 -0
- scripts/fake_data.py +4 -0
- scripts/inference.py +10 -0
- scripts/result.py +4 -0
- scripts/train.py +13 -0
- weight/.gitkeep +0 -0
.gitattributes
CHANGED
|
@@ -1,35 +1,2 @@
|
|
| 1 |
-
*.
|
| 2 |
-
*.
|
| 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 binary
|
| 2 |
+
*.npz binary
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
conf/config.yaml
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
seed: 42
|
| 2 |
+
data: {path: data/downscale.npz, samples: 6, variables: 4, window: 3, coarse_grid: [8, 8], fine_grid: [128, 128]}
|
| 3 |
+
model: {channels: 12, hidden: 16}
|
| 4 |
+
train: {epochs: 1, learning_rate: 0.001}
|
| 5 |
+
inference: {members: 8, denoise_steps: 4}
|
| 6 |
+
paths: {checkpoint: result/checkpoints/climate2weather.pt, training_metrics: result/training/metrics.json, predictions: result/output/predictions.npz, evaluation: result/evaluation/metrics.json, figure: result/evaluation/comparison.png}
|
| 7 |
+
paper_model: {variables: [uas, vas, tas, psl], coarse_grid: [8, 8], fine_grid: [128, 128], coarse_hours: 6, fine_hours: 1, train_years: [2006, 2013], test_year: 2014}
|
config.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{"model_name":"Climate2Weather","model_type":"climate2weather","architectures":["ScoreUNet"],"framework":"PyTorch","domain":"climate-downscaling","task":"probabilistic-spatiotemporal-downscaling","implementation":{"entry_point":"model/climate2weather.py","scope":"L2 score-based data-assimilation reproduction"},"architecture":{"variables":["uas","vas","tas","psl"],"coarse_grid":[8,8],"fine_grid":[128,128],"coarse_hours":6,"fine_hours":1,"engineering_window":3,"paper_parameters":72000000},"data":{"target":"COSMO-REA6","conditions":["MPI-ESM1.2-HR","HadGEM3-GC3.1-LM"],"train_years":[2006,2013],"test_year":2014},"configuration_sources":["conf/config.yaml","model/climate2weather.py","scripts/fake_data.py","scripts/train.py","scripts/inference.py","scripts/result.py"]}
|
model/climate2weather.py
ADDED
|
@@ -0,0 +1,15 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import json
|
| 2 |
+
from pathlib import Path
|
| 3 |
+
import numpy as np
|
| 4 |
+
import torch
|
| 5 |
+
from torch import nn
|
| 6 |
+
import torch.nn.functional as F
|
| 7 |
+
import yaml
|
| 8 |
+
def cfg(r):return yaml.safe_load((Path(r)/"conf/config.yaml").read_text())
|
| 9 |
+
def truth(i):
|
| 10 |
+
y,x=torch.meshgrid(torch.linspace(-1,1,128),torch.linspace(-1,1,128),indexing='ij');return torch.stack([torch.sin((v+1)*x+.1*i)*torch.cos((v+1)*y) for v in range(12)]).float()
|
| 11 |
+
def observe(x):return F.avg_pool2d(x.reshape(3,4,128,128),16).reshape(12,8,8)
|
| 12 |
+
class ScoreUNet(nn.Module):
|
| 13 |
+
def __init__(self,channels=12,hidden=16):super().__init__();self.net=nn.Sequential(nn.Conv2d(channels+1,hidden,3,padding=1),nn.SiLU(),nn.Conv2d(hidden,hidden,3,padding=1),nn.SiLU(),nn.Conv2d(hidden,channels,3,padding=1));self.model_config={'channels':channels,'hidden':hidden}
|
| 14 |
+
def forward(self,x,sigma):return self.net(torch.cat((x,torch.ones_like(x[:,:1])*sigma),1))
|
| 15 |
+
def write(p,o):p=Path(p);p.parent.mkdir(parents=True,exist_ok=True);p.write_text(json.dumps(o,indent=2)+'\n')
|
scripts/fake_data.py
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from pathlib import Path
|
| 2 |
+
import sys,numpy as np
|
| 3 |
+
R=Path(__file__).resolve().parents[1];sys.path.insert(0,str(R));from model.climate2weather import *
|
| 4 |
+
c=cfg(R);f=np.stack([truth(i).numpy() for i in range(6)]);co=np.stack([observe(torch.tensor(x)).numpy() for x in f]);p=R/c['data']['path'];p.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(p,fine=f,coarse=co,variables=['uas','vas','tas','psl'],fine_shape=[3,4,128,128],coarse_shape=[3,4,8,8]);print(p)
|
scripts/inference.py
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from pathlib import Path
|
| 2 |
+
import sys,numpy as np,torch
|
| 3 |
+
R=Path(__file__).resolve().parents[1];sys.path.insert(0,str(R));from model.climate2weather import *
|
| 4 |
+
c=cfg(R);d=np.load(R/c['data']['path']);z=torch.load(R/c['paths']['checkpoint'],map_location='cpu',weights_only=True);m=ScoreUNet(**z['model_config']);m.load_state_dict(z['model']);condition=torch.tensor(d['coarse'][-1:]);up=F.interpolate(condition,size=(128,128),mode='nearest');ens=[]
|
| 5 |
+
with torch.no_grad():
|
| 6 |
+
for j in range(8):
|
| 7 |
+
x=up+.2*torch.randn_like(up)
|
| 8 |
+
for k in range(4):x=x+.02*m(x,.2/(k+1));x=x+(up-F.interpolate(observe(x[0])[None],size=(128,128),mode='nearest'))*.2
|
| 9 |
+
ens.append(x[0].reshape(3,4,128,128).numpy())
|
| 10 |
+
p=R/c['paths']['predictions'];p.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(p,ensemble=ens,target=d['fine'][-1].reshape(3,4,128,128),condition=condition.numpy().reshape(1,3,4,8,8));print(p)
|
scripts/result.py
ADDED
|
@@ -0,0 +1,4 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from pathlib import Path
|
| 2 |
+
import sys,numpy as np;import matplotlib;matplotlib.use('Agg');import matplotlib.pyplot as plt
|
| 3 |
+
R=Path(__file__).resolve().parents[1];sys.path.insert(0,str(R));from model.climate2weather import cfg,write
|
| 4 |
+
c=cfg(R);d=np.load(R/c['paths']['predictions']);e=d['ensemble'];t=d['target'];mean=e.mean(0);rmse=float(np.sqrt(np.mean((mean-t)**2)));spread=float(e.std(0).mean());pit=float(np.mean(e<t[None]));coherence=float(np.mean(abs(np.diff(mean,axis=0))));write(R/c['paths']['evaluation'],{'rmse':rmse,'spread':spread,'pit_mean':pit,'temporal_difference':coherence,'synthetic':True});plt.imshow(mean[0,0]-t[0,0],cmap='coolwarm');plt.colorbar();p=R/c['paths']['figure'];p.parent.mkdir(parents=True,exist_ok=True);plt.savefig(p,dpi=150);print(p)
|
scripts/train.py
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
from pathlib import Path
|
| 2 |
+
import sys,os,numpy as np,torch;import torch.distributed as dist
|
| 3 |
+
from torch.nn.parallel import DistributedDataParallel as DDP
|
| 4 |
+
R=Path(__file__).resolve().parents[1];sys.path.insert(0,str(R));from model.climate2weather import *
|
| 5 |
+
c=cfg(R);rank=int(os.getenv('RANK',0));world=int(os.getenv('WORLD_SIZE',1));ddp=world>1
|
| 6 |
+
if ddp:dist.init_process_group('gloo')
|
| 7 |
+
d=np.load(R/c['data']['path']);base=ScoreUNet(**c['model']);m=DDP(base) if ddp else base;opt=torch.optim.Adam(m.parameters(),lr=c['train']['learning_rate']);ls=[]
|
| 8 |
+
for i in range(rank,len(d['fine']),world):x=torch.tensor(d['fine'][i:i+1]);noise=torch.randn_like(x);s=.2;z=x+s*noise;loss=((m(z,s)+noise/s)**2).mean();opt.zero_grad();loss.backward();opt.step();ls.append(float(loss))
|
| 9 |
+
v=torch.tensor([sum(ls),len(ls)],dtype=torch.float64)
|
| 10 |
+
if ddp:dist.all_reduce(v)
|
| 11 |
+
p=R/c['paths']['checkpoint']
|
| 12 |
+
if rank==0:p.parent.mkdir(parents=True,exist_ok=True);torch.save({'model':base.state_dict(),'model_config':c['model']},p);write(R/c['paths']['training_metrics'],{'score_loss':float(v[0]/v[1]),'world_size':world});print(p)
|
| 13 |
+
if ddp:dist.destroy_process_group()
|
weight/.gitkeep
ADDED
|
File without changes
|