TropiCycloneNet / model /tropicyclonenet.py
zhangrenchao's picture
Publish TropiCycloneNet reproduction
01384b4 verified
Raw History Blame Contribute Delete
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")