zhangrenchao commited on
Commit
21b9cf1
·
verified ·
1 Parent(s): 68fd2c6

Publish Climate2Weather 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 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