| """WeatherBench five-layer fully convolutional baseline.""" |
| import json |
| from pathlib import Path |
| import numpy as np |
| import torch |
| from torch import nn |
| import torch.nn.functional as F |
| import yaml |
|
|
| PRESSURE_LEVELS=(50,100,150,200,250,300,400,500,600,700,850,925,1000) |
| def load_config(root):return yaml.safe_load((Path(root)/"conf/config.yaml").read_text()) |
| def periodic_conv(x,conv,pad=2):x=F.pad(x,(pad,pad,0,0),mode="circular");x=F.pad(x,(0,0,pad,pad),mode="replicate");return conv(x) |
| class WeatherBenchCNN(nn.Module): |
| def __init__(self,hidden_channels=16,layers=5,kernel_size=5): |
| super().__init__();chs=[2]+[hidden_channels]*(layers-1)+[2];self.convs=nn.ModuleList([nn.Conv2d(chs[i],chs[i+1],kernel_size) for i in range(layers)]);self.model_config={"hidden_channels":hidden_channels,"layers":layers,"kernel_size":kernel_size} |
| def forward(self,x): |
| if x.shape[1:]!=(2,32,64):raise ValueError("expected [B,2,32,64]") |
| for c in self.convs[:-1]:x=F.elu(periodic_conv(x,c)) |
| return periodic_conv(x,self.convs[-1]) |
| def synthetic_state(i,lead=0): |
| lat,lon=torch.meshgrid(torch.linspace(-90,90,32),torch.arange(64).float()*360/64,indexing="ij");phase=.12*(i+lead);z=50000+3000*torch.cos(torch.deg2rad(lat))*torch.sin(torch.deg2rad(lon)+phase);t=270-35*abs(lat)/90+4*torch.cos(torch.deg2rad(lon*2)-phase);return torch.stack((z/50000,t/270)).float() |
| def weighted_rmse(p,t,lat):w=np.cos(np.deg2rad(lat));w=w/w.mean();return np.mean(np.sqrt(np.mean((p-t)**2*w[None,None,:,None],axis=(1,2,3)))) |
| def weighted_acc(p,t,clim,lat):w=np.cos(np.deg2rad(lat))[None,None,:,None];a=p-clim;b=t-clim;return float((w*a*b).sum()/np.sqrt((w*a*a).sum()*(w*b*b).sum())) |
| 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") |
|
|