File size: 2,884 Bytes
01384b4 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 | """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")
|