typhoon-predict / run_inference.py
euler314's picture
Add trained checkpoint and inference artifacts
dfe8a5e verified
Raw
History Blame Contribute Delete
4.96 kB
import json
import argparse
from pathlib import Path
import numpy as np
import torch
from torch import nn
ROOT = Path(__file__).resolve().parent
parser = argparse.ArgumentParser(description="Run the Bavi typhoon checkpoint locally.")
parser.add_argument("--checkpoint", type=Path, default=ROOT / "best.pt")
parser.add_argument("--output", type=Path, default=ROOT / "forecast.json")
parser.add_argument("--device", choices=["auto", "mps", "cpu"], default="auto")
args = parser.parse_args()
CKPT = args.checkpoint.expanduser().resolve()
OUT = args.output.expanduser().resolve()
class FieldEncoder(nn.Module):
def __init__(self, channels, hidden=128):
super().__init__()
self.net = nn.Sequential(
nn.Conv2d(channels,64,3,padding=1), nn.GELU(), nn.BatchNorm2d(64),
nn.Conv2d(64,96,3,stride=2,padding=1), nn.GELU(), nn.BatchNorm2d(96),
nn.Conv2d(96,hidden,3,stride=2,padding=1), nn.GELU(),
nn.AdaptiveAvgPool2d(1))
def forward(self, x): return self.net(x).flatten(1)
class ERA5CycloneEnsemble(nn.Module):
def __init__(self, track_dim, channels, steps, output_dim, hidden=192, latent=32):
super().__init__()
self.steps, self.output_dim, self.latent = steps, output_dim, latent
self.field = FieldEncoder(channels, 128)
self.track_gru = nn.GRU(track_dim,96,batch_first=True,bidirectional=True)
self.fuse = nn.Sequential(nn.Linear(128*steps+192,hidden),nn.GELU(),nn.Dropout(.15),nn.Linear(hidden,hidden),nn.GELU())
self.mean = nn.Linear(hidden,output_dim)
self.log_scale = nn.Linear(hidden,output_dim)
self.latent_proj = nn.Sequential(nn.Linear(latent,hidden),nn.GELU(),nn.Linear(hidden,output_dim))
def encode(self, track_x, field_x):
b,t,c,h,w=field_x.shape
field_z=self.field(field_x.reshape(b*t,c,h,w)).reshape(b,t,-1).flatten(1)
_,state=self.track_gru(track_x)
return self.fuse(torch.cat([field_z,state.transpose(0,1).flatten(1)],dim=1))
def sample(self, track_x, field_x, n, temperature=1.0):
z=self.encode(track_x,field_x)
mean=self.mean(z); log_scale=self.log_scale(z).clamp(-5,2)
eps=torch.randn(n,z.shape[0],self.latent,device=z.device)
low=self.latent_proj(eps.reshape(-1,self.latent)).reshape(n,z.shape[0],-1)
independent=torch.randn_like(low)*torch.exp(log_scale).unsqueeze(0)
return mean.unsqueeze(0)+temperature*(low+independent)
raw=torch.load(CKPT,map_location="cpu",weights_only=False)
if args.device == "mps":
device = torch.device("mps")
elif args.device == "cpu":
device = torch.device("cpu")
else:
device=torch.device("mps" if torch.backends.mps.is_available() else "cpu")
config=raw["config"]
track_scaler=raw["track_scaler"]; y_scaler=raw["y_scaler"]
field_mean=np.asarray(raw["field_mean"],dtype="float32")
field_std=np.asarray(raw["field_std"],dtype="float32")
model=ERA5CycloneEnsemble(9,10,1,28).to(device)
model.load_state_dict(raw["model_state"]); model.eval()
fixes=[
("2026-07-09T12:00:00",19.2,128.8,100,952),
("2026-07-09T18:00:00",20.1,128.2,90,953),
("2026-07-10T00:00:00",20.8,127.3,75,964),
("2026-07-10T06:00:00",21.9,126.9,75,962),
]
track=[]
for i,(stamp,lat,lon,wind,pres) in enumerate(fixes):
dlat=0 if i==0 else lat-fixes[i-1][1]
dlon=0 if i==0 else ((lon-fixes[i-1][2]+180)%360)-180
track.append([lat,lon,wind,pres,dlat,dlon,float(np.hypot(dlat,dlon)),np.sin(2*np.pi*191/366),np.cos(2*np.pi*191/366)])
track=np.asarray(track,dtype="float32")
track[:,:1]-=track[-1:,0:1]
track[:,1:2]=((track[:,1:2]-track[-1:,1:2]+180)%360)-180
track[:,2]/=100.0
track[:,3]=(track[:,3]-950.0)/50.0
xtrack=track_scaler.transform(track).astype("float32")[None]
field=np.zeros((1,1,10,33,33),dtype="float32")
with torch.no_grad():
ens=model.sample(torch.from_numpy(xtrack).to(device),torch.from_numpy(field).to(device),50,float(config.get("sample_temperature",1.0))).cpu().numpy()
ens=ens.mean(0,keepdims=True)+(ens-ens.mean(0,keepdims=True))*1.1297996044158936
pred=y_scaler.inverse_transform(ens[:,0,:])
base_lat,base_lon=fixes[-1][1],fixes[-1][2]
points=[]
for k,lead in enumerate(config["lead_hours"]):
j=4*k
lat=base_lat+pred[:,j]; lon=(base_lon+pred[:,j+1])%360
points.append({"lead_hours":int(lead),"lat":float(lat.mean()),"lon":float(lon.mean()),"p10_lat":float(np.quantile(lat,.1)),"p90_lat":float(np.quantile(lat,.9)),"p10_lon":float(np.quantile(lon,.1)),"p90_lon":float(np.quantile(lon,.9))})
result={"storm":"Bavi","source":"JTWC/TCGP fixes through 2026-07-10 0600 UTC","initial_time":"2026-07-10T06:00:00Z","initial_lat":base_lat,"initial_lon":base_lon,"points":points,"device":str(device),"note":"Mac checkpoint inference using current track fixes and mean-normalized atmospheric input because 2026 ERA5 fields are unavailable locally. Not an operational forecast."}
OUT.write_text(json.dumps(result,indent=2))
print(json.dumps(result,indent=2))