zhangrenchao commited on
Commit
01384b4
·
verified ·
1 Parent(s): ae2d9f8

Publish TropiCycloneNet 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 filter=lfs diff=lfs merge=lfs -text
2
+ *.npz filter=lfs diff=lfs merge=lfs -text
 
 
 
 
 
 
 
 
 
 
 
 
conf/config.yaml ADDED
@@ -0,0 +1,35 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ seed: 42
2
+ data:
3
+ format_version: tcnd_synthetic_v1
4
+ manifest: data/manifest.json
5
+ samples: 8
6
+ history_steps: 8
7
+ forecast_steps: 4
8
+ interval_hours: 6
9
+ data1d_features: 4
10
+ data3d_channels: 1
11
+ grid: [81, 81]
12
+ environment_features: 73
13
+ basins: [NA, EP, WP, NI, SI, SP]
14
+ model:
15
+ hidden_dim: 24
16
+ generators: 6
17
+ heads: 4
18
+ train:
19
+ epochs: 1
20
+ batch_size: 2
21
+ learning_rate: 0.001
22
+ paths:
23
+ checkpoint: result/checkpoints/tropicyclonenet.pt
24
+ training_metrics: result/training/metrics.json
25
+ predictions: result/output/predictions.npz
26
+ evaluation: result/evaluation/metrics.json
27
+ figure: result/evaluation/comparison.png
28
+ paper_model:
29
+ training_years: [1950, 2016]
30
+ test_years: [2017, 2021]
31
+ optimizer: Adam
32
+ learning_rate: 0.0001
33
+ batch_size: 96
34
+ epochs: 102
35
+ generators: 6
config.json ADDED
@@ -0,0 +1 @@
 
 
1
+ {"model_name":"TropiCycloneNet","model_type":"tropicyclonenet","architectures":["TropiCycloneNet"],"framework":"PyTorch","domain":"tropical-cyclone","task":"multimodal-track-intensity-forecasting","implementation":{"entry_point":"model/tropicyclonenet.py","scope":"core-method engineering reproduction"},"architecture":{"history_steps":8,"forecast_steps":4,"interval_hours":6,"data1d_variables":["longitude","latitude","pressure","wind"],"data3d_shape":[1,81,81],"environment_features":73,"paper_generators":6,"core":["3D data encoder","1D LSTM encoder","environment-time attention","generator chooser","multiple decoder LSTMs"]},"data":{"name":"TCND","years":[1950,2021],"cyclones":3630,"basins":6,"spatial_resolution_degrees":0.25,"synthetic":true},"configuration_sources":["conf/config.yaml","model/tropicyclonenet.py","scripts/fake_data.py","scripts/train.py","scripts/inference.py","scripts/result.py"]}
model/tropicyclonenet.py ADDED
@@ -0,0 +1,30 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Compact multimodal TropiCycloneNet reproduction."""
2
+ import json,math
3
+ from pathlib import Path
4
+ import numpy as np
5
+ import torch
6
+ from torch import nn
7
+ import torch.nn.functional as F
8
+ import yaml
9
+
10
+ def load_config(root):return yaml.safe_load((Path(root)/"conf/config.yaml").read_text())
11
+ def synthetic_sample(i):
12
+ torch.manual_seed(1000+i);t=torch.arange(12).float();lon=120+i*.7+1.4*t+.08*t*t;lat=12+i*.3+.65*t;pressure=995-2.4*t-i*.2;wind=18+1.7*t+i*.1;one=torch.stack((lon/180,lat/90,(pressure-960)/50,(wind-40)/25),1);yy,xx=torch.meshgrid(torch.linspace(-1,1,81),torch.linspace(-1,1,81),indexing="ij");grids=[]
13
+ for k in range(12):grids.append(torch.exp(-((xx-.03*k)**2+(yy+.02*k)**2)/.15)[None])
14
+ env=torch.zeros(8,73);env[:,0]=1.5;env[:,1+(i%12)]=1;env[:,13+(i%36)]=1;env[:,49+(i%12)]=1;env[:,61+(i%8)]=1;env[:,69+(i%4)]=1
15
+ return one[:8],torch.stack(grids[:8]),env,one[8:],torch.stack(grids[8:])
16
+
17
+ class TropiCycloneNet(nn.Module):
18
+ def __init__(self,hidden_dim=24,generators=6,heads=4):
19
+ super().__init__();self.generators=generators;self.grid_encoder=nn.Sequential(nn.Conv3d(1,8,3,padding=1),nn.GELU(),nn.AdaptiveAvgPool3d((8,4,4)));self.grid_project=nn.Linear(8*4*4,hidden_dim);self.one_encoder=nn.LSTM(4+hidden_dim,hidden_dim,batch_first=True);self.env_project=nn.Linear(73,hidden_dim);layer=nn.TransformerEncoderLayer(hidden_dim,heads,hidden_dim*2,batch_first=True);self.env_time=nn.TransformerEncoder(layer,1);self.chooser=nn.Linear(hidden_dim*2,generators);self.decoders=nn.ModuleList([nn.LSTMCell(4+hidden_dim*2,hidden_dim) for _ in range(generators)]);self.heads=nn.ModuleList([nn.Linear(hidden_dim,4) for _ in range(generators)]);self.model_config={"hidden_dim":hidden_dim,"generators":generators,"heads":heads}
20
+ def forward(self,one,grid,env,steps=4):
21
+ if one.shape[1:]!=(8,4) or grid.shape[1:]!=(8,1,81,81) or env.shape[1:]!=(8,73):raise ValueError("invalid TCND shape")
22
+ g=self.grid_encoder(grid.transpose(1,2)).permute(0,2,1,3,4).flatten(2);g=self.grid_project(g);_,(h,_)=self.one_encoder(torch.cat((one,g),-1));e=self.env_time(self.env_project(env))[:,-1];context=torch.cat((h[-1],e),-1);prob=self.chooser(context).softmax(-1);all_outputs=[]
23
+ for decoder,head in zip(self.decoders,self.heads):
24
+ hx=h[-1];cx=torch.zeros_like(hx);prev=one[:,-1];seq=[]
25
+ for _ in range(steps):hx,cx=decoder(torch.cat((prev,context),-1),(hx,cx));prev=prev+head(hx);seq.append(prev)
26
+ all_outputs.append(torch.stack(seq,1))
27
+ return torch.stack(all_outputs,1),prob
28
+ def haversine_km(a,b):
29
+ lon1,lat1,lon2,lat2=[torch.deg2rad(x) for x in (a[...,0]*180,a[...,1]*90,b[...,0]*180,b[...,1]*90)];dlat=lat2-lat1;dlon=lon2-lon1;q=torch.sin(dlat/2)**2+torch.cos(lat1)*torch.cos(lat2)*torch.sin(dlon/2)**2;return 6371*2*torch.asin(torch.sqrt(q.clamp(0,1)))
30
+ def write_json(path,obj):path=Path(path);path.parent.mkdir(parents=True,exist_ok=True);path.write_text(json.dumps(obj,indent=2)+"\n")
scripts/fake_data.py ADDED
@@ -0,0 +1,5 @@
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ import sys,json
3
+ ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
4
+ from model.tropicyclonenet import *
5
+ c=load_config(ROOT);o,g,e,y,gy=synthetic_sample(0);assert o.shape==(8,4) and g.shape==(8,1,81,81) and e.shape==(8,73) and y.shape==(4,4);p=ROOT/c["data"]["manifest"];p.parent.mkdir(parents=True,exist_ok=True);p.write_text(json.dumps({"format_version":c["data"]["format_version"],"data1d_history_shape":[8,4],"data3d_history_shape":[8,1,81,81],"environment_shape":[8,73],"target_shape":[4,4],"interval_hours":6,"basins":c["data"]["basins"],"synthetic":True},indent=2));print(p)
scripts/inference.py ADDED
@@ -0,0 +1,8 @@
 
 
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ import sys,numpy as np,torch
3
+ ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
4
+ from model.tropicyclonenet import *
5
+ c=load_config(ROOT);ck=torch.load(ROOT/c["paths"]["checkpoint"],map_location="cpu",weights_only=True);m=TropiCycloneNet(**ck["model_config"]);m.load_state_dict(ck["model"]);m.eval();pred=[];truth=[];prob=[]
6
+ with torch.no_grad():
7
+ for i in range(c["data"]["samples"]):o,g,e,y,_=synthetic_sample(100+i);p,q=m(o[None],g[None],e[None]);pred.append(p[0].numpy());prob.append(q[0].numpy());truth.append(y.numpy())
8
+ path=ROOT/c["paths"]["predictions"];path.parent.mkdir(parents=True,exist_ok=True);np.savez_compressed(path,prediction=np.asarray(pred),target=np.asarray(truth),generator_probability=np.asarray(prob),lead_hours=np.arange(1,5)*6);print(path)
scripts/result.py ADDED
@@ -0,0 +1,6 @@
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ import sys,numpy as np,torch
3
+ import matplotlib;matplotlib.use("Agg");import matplotlib.pyplot as plt
4
+ ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
5
+ from model.tropicyclonenet import *
6
+ c=load_config(ROOT);d=np.load(ROOT/c["paths"]["predictions"]);p=torch.from_numpy(d["prediction"]);t=torch.from_numpy(d["target"]);best=((p-t[:,None])**2).mean((2,3)).argmin(1);chosen=p[torch.arange(len(p)),best];track=haversine_km(chosen[...,:2],t[...,:2]).mean(0);pres=(abs(chosen[...,2]-t[...,2])*50).mean(0);wind=(abs(chosen[...,3]-t[...,3])*25).mean(0);write_json(ROOT/c["paths"]["evaluation"],{"track_mae_km":track.tolist(),"pressure_mae_hpa":pres.tolist(),"wind_mae_ms":wind.tolist(),"generators":int(p.shape[1]),"synthetic":True});fig,ax=plt.subplots(1,2,figsize=(9,3.5));ax[0].plot(d["lead_hours"],track,"o-");ax[0].set(xlabel="Lead (h)",ylabel="Track MAE (km)");ax[1].plot(t[0,:,0]*180,t[0,:,1]*90,"ko-",label="truth");ax[1].plot(chosen[0,:,0]*180,chosen[0,:,1]*90,"r.--",label="prediction");ax[1].legend();ax[1].set(xlabel="Longitude",ylabel="Latitude");fig.tight_layout();path=ROOT/c["paths"]["figure"];path.parent.mkdir(parents=True,exist_ok=True);fig.savefig(path,dpi=150);print(path)
scripts/train.py ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from pathlib import Path
2
+ import sys,os,math,torch
3
+ import torch.distributed as dist
4
+ from torch.nn.parallel import DistributedDataParallel as DDP
5
+ ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
6
+ from model.tropicyclonenet import *
7
+ c=load_config(ROOT);rank=int(os.getenv("RANK",0));world=int(os.getenv("WORLD_SIZE",1));distributed=world>1
8
+ if distributed:dist.init_process_group("gloo")
9
+ torch.manual_seed(c["seed"]);base=TropiCycloneNet(**c["model"]);m=DDP(base) if distributed else base;opt=torch.optim.Adam(m.parameters(),lr=c["train"]["learning_rate"]);idx=list(range(rank,c["data"]["samples"],world));losses=[]
10
+ for _ in range(c["train"]["epochs"]):
11
+ for i in idx:o,g,e,y,_=synthetic_sample(i);pred,prob=m(o[None],g[None],e[None]);errors=((pred-y[None,None])**2).mean((2,3));best=errors.min(1).values.mean();diversity=-pred.std(1).mean();loss=best+.01*diversity+.01*(-torch.log(prob+1e-8).mean());opt.zero_grad();loss.backward();opt.step();losses.append(float(loss))
12
+ total=torch.tensor([sum(losses),len(losses)],dtype=torch.float64)
13
+ if distributed:dist.all_reduce(total)
14
+ p=ROOT/c["paths"]["checkpoint"]
15
+ if rank==0:p.parent.mkdir(parents=True,exist_ok=True);torch.save({"model":base.state_dict(),"model_config":c["model"]},p);write_json(ROOT/c["paths"]["training_metrics"],{"loss":float(total[0]/total[1]),"world_size":world});print(p)
16
+ if distributed:dist.destroy_process_group()
weight/.gitkeep ADDED
File without changes