Download model/tropicyclonenet.py from OneScience-Group/TropiCycloneNet: direct link, hf CLI and curl.
- Browser
- Download file 2.88 kB
-
https://huggingface.co/OneScience-Group/TropiCycloneNet/resolve/main/model/tropicyclonenet.py
- Command line
-
hf download hf://OneScience-Group/TropiCycloneNet/model/tropicyclonenet.py
-
curl -L -o tropicyclonenet.py https://huggingface.co/OneScience-Group/TropiCycloneNet/resolve/main/model/tropicyclonenet.py
2.88 kB
| """Compact multimodal TropiCycloneNet reproduction.""" | |
| import json,math | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| from torch import nn | |
| import torch.nn.functional as F | |
| import yaml | |
| def load_config(root):return yaml.safe_load((Path(root)/"conf/config.yaml").read_text()) | |
| def synthetic_sample(i): | |
| 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=[] | |
| for k in range(12):grids.append(torch.exp(-((xx-.03*k)**2+(yy+.02*k)**2)/.15)[None]) | |
| 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 | |
| return one[:8],torch.stack(grids[:8]),env,one[8:],torch.stack(grids[8:]) | |
| class TropiCycloneNet(nn.Module): | |
| def __init__(self,hidden_dim=24,generators=6,heads=4): | |
| 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} | |
| def forward(self,one,grid,env,steps=4): | |
| 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") | |
| 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=[] | |
| for decoder,head in zip(self.decoders,self.heads): | |
| hx=h[-1];cx=torch.zeros_like(hx);prev=one[:,-1];seq=[] | |
| for _ in range(steps):hx,cx=decoder(torch.cat((prev,context),-1),(hx,cx));prev=prev+head(hx);seq.append(prev) | |
| all_outputs.append(torch.stack(seq,1)) | |
| return torch.stack(all_outputs,1),prob | |
| def haversine_km(a,b): | |
| 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))) | |
| 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") | |