Publish DLESyM reproduction
Browse files- .gitattributes +2 -35
- conf/config.yaml +7 -0
- config.json +1 -0
- model/dlesym.py +20 -0
- scripts/fake_data.py +6 -0
- scripts/inference.py +7 -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/dlesym.npz, samples: 6, tile: 16, atmosphere_channels: 9, logical_grid: [180, 360]}
|
| 3 |
+
model: {atmosphere_channels: 9, hidden: 16}
|
| 4 |
+
train: {epochs: 1, learning_rate: 0.001}
|
| 5 |
+
inference: {cycles: 4}
|
| 6 |
+
paths: {checkpoint: result/checkpoints/dlesym.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: {resolution_km: 110, atmosphere_channels: 9, ocean_channels: 1, atmosphere_step_hours: 6, ocean_step_hours: 48, coupling_cycle_hours: 96, atmosphere_epochs: 250, ocean_epochs: 300}
|
config.json
ADDED
|
@@ -0,0 +1 @@
|
|
|
|
|
|
|
| 1 |
+
{"model_name":"DLESyM","model_type":"dlesym","architectures":["DLESyM"],"framework":"PyTorch","domain":"coupled-climate","task":"earth-system-simulation","implementation":{"entry_point":"model/dlesym.py","scope":"L2 coupled core reproduction"},"architecture":{"atmosphere_fields":9,"ocean_fields":["SST"],"atmosphere_step_hours":6,"ocean_step_hours":48,"coupling_cycle_hours":96,"logical_grid":[180,360],"modules":["DLWP","DLOM","precipitation"]},"data":{"sources":["ERA5","ISCCP"],"train_years":[1983,2016],"synthetic_tiles":true},"configuration_sources":["conf/config.yaml","model/dlesym.py","scripts/fake_data.py","scripts/train.py","scripts/inference.py","scripts/result.py"]}
|
model/dlesym.py
ADDED
|
@@ -0,0 +1,20 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
import json,torch
|
| 2 |
+
from pathlib import Path
|
| 3 |
+
from torch import nn
|
| 4 |
+
import yaml
|
| 5 |
+
FIELDS=('Z1000','Z500','Z250','THICK300_700','T2M','T850','TCWV','WS10','OLR')
|
| 6 |
+
def cfg(r):return yaml.safe_load((Path(r)/'conf/config.yaml').read_text())
|
| 7 |
+
def state(i,s=0):
|
| 8 |
+
y,x=torch.meshgrid(torch.linspace(-1,1,16),torch.linspace(-1,1,16),indexing='ij');c=torch.arange(9)[:,None,None];a=torch.sin((c+1)*x+.04*(i+s))*torch.cos((c%4+1)*y);sst=torch.cos(x+.01*(i+s))*torch.cos(y);return a.float(),sst[None].float()
|
| 9 |
+
class Net(nn.Module):
|
| 10 |
+
def __init__(self,ci,co,h):super().__init__();self.n=nn.Sequential(nn.Conv2d(ci,h,3,padding=1),nn.GELU(),nn.Conv2d(h,h,3,padding=1),nn.GELU(),nn.Conv2d(h,co,1))
|
| 11 |
+
def forward(self,x):return self.n(x)
|
| 12 |
+
class DLESyM(nn.Module):
|
| 13 |
+
def __init__(self,atmosphere_channels=9,hidden=16):super().__init__();self.atm=Net(10,9,hidden);self.ocean=Net(4,1,hidden//2);self.precip=Net(9,1,hidden//2);self.model_config={'atmosphere_channels':atmosphere_channels,'hidden':hidden}
|
| 14 |
+
def atmosphere(self,a,sst):return a+self.atm(torch.cat((a,sst),1))
|
| 15 |
+
def cycle(self,a,sst):
|
| 16 |
+
seq=[]
|
| 17 |
+
for _ in range(16):a=self.atmosphere(a,sst);seq.append(a)
|
| 18 |
+
f=torch.stack(seq);forcing=torch.stack((f[:,:,7].mean(0),f[:,:,0].mean(0),f[:,:,8].mean(0)),1);sst=sst+self.ocean(torch.cat((sst,forcing),1));return a,sst,self.precip(a)
|
| 19 |
+
def forward(self,a,sst):return self.cycle(a,sst)
|
| 20 |
+
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,6 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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.dlesym import *
|
| 4 |
+
c=cfg(R);a=[];s=[];ta=[];ts=[]
|
| 5 |
+
for i in range(6):x,y=state(i);u,v=state(i,1);a.append(x);s.append(y);ta.append(u);ts.append(v)
|
| 6 |
+
p=R/c['data']['path'];p.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(p,atmosphere=a,sst=s,target_atmosphere=ta,target_sst=ts,fields=FIELDS,logical_shape=[10,180,360],is_complete_global=False);print(p)
|
scripts/inference.py
ADDED
|
@@ -0,0 +1,7 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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.dlesym import *
|
| 4 |
+
c=cfg(R);z=torch.load(R/c['paths']['checkpoint'],map_location='cpu',weights_only=True);m=DLESyM(**z['model_config']);m.load_state_dict(z['model']);a,s=state(20);a=a[None];s=s[None];aa=[];ss=[];pp=[]
|
| 5 |
+
with torch.no_grad():
|
| 6 |
+
for _ in range(4):a,s,p=m(a,s);aa.append(a[0].numpy());ss.append(s[0].numpy());pp.append(p[0].numpy())
|
| 7 |
+
q=R/c['paths']['predictions'];q.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(q,atmosphere=aa,sst=ss,precipitation=pp,lead_days=np.arange(1,5)*4,logical_shape=[10,180,360],is_complete_global=False);print(q)
|
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.dlesym import cfg,write
|
| 4 |
+
c=cfg(R);d=np.load(R/c['paths']['predictions']);a=d['atmosphere'];s=d['sst'];drift=np.sqrt(np.mean((a-a[:1])**2,axis=(1,2,3)));sst=s.mean((1,2,3));prec=d['precipitation'].mean((1,2,3));write(R/c['paths']['evaluation'],{'atmosphere_drift_rmse':drift.tolist(),'sst_mean':sst.tolist(),'precipitation_mean':prec.tolist(),'is_complete_global':False,'synthetic':True});plt.plot(d['lead_days'],drift);plt.xlabel('Lead (days)');plt.ylabel('Atmosphere drift RMSE');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.dlesym 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=DLESyM(**c['model']);m=DDP(base) if ddp else base;opt=torch.optim.AdamW(m.parameters(),lr=c['train']['learning_rate']);ls=[]
|
| 8 |
+
for i in range(rank,6,world):a,s=map(torch.tensor,(d['atmosphere'][i:i+1],d['sst'][i:i+1]));pa,ps,pr=m(a,s);loss=((pa-torch.tensor(d['target_atmosphere'][i:i+1]))**2).mean()+((ps-torch.tensor(d['target_sst'][i:i+1]))**2).mean()+pr.pow(2).mean()*.01;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'],{'loss':float(v[0]/v[1]),'world_size':world});print(p)
|
| 13 |
+
if ddp:dist.destroy_process_group()
|
weight/.gitkeep
ADDED
|
File without changes
|