Publish TropiCycloneNet reproduction
Browse files- .gitattributes +1 -34
- conf/config.yaml +35 -0
- config.json +1 -0
- model/tropicyclonenet.py +30 -0
- scripts/fake_data.py +5 -0
- scripts/inference.py +8 -0
- scripts/result.py +6 -0
- scripts/train.py +16 -0
- weight/.gitkeep +0 -0
.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 |
-
*.
|
| 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
|