import json,torch from pathlib import Path from torch import nn import torch.nn.functional as F import yaml def cfg(r):return yaml.safe_load((Path(r)/'conf/config.yaml').read_text()) def sample(i): y,x=torch.meshgrid(torch.linspace(-1,1,96),torch.linspace(-1,1,96),indexing='ij');tc=((x-.3*torch.sin(torch.tensor(i)))**2+(y-.2)**2<.08).long();ar=(abs(y-.4*x)<.08).long()*2;lab=torch.maximum(tc,ar);f=torch.stack((torch.exp(-((x)**2+y**2)),torch.sin(x*4),torch.cos(y*4),(lab>0).float()));return f.float(),lab class ClimateNetDeepLab(nn.Module): def __init__(self,channels=4,classes=3,hidden=16):super().__init__();self.enc=nn.Sequential(nn.Conv2d(channels,hidden,3,padding=1),nn.BatchNorm2d(hidden),nn.ReLU(),nn.Conv2d(hidden,hidden*2,3,2,1),nn.ReLU());self.aspp=nn.ModuleList([nn.Conv2d(hidden*2,hidden,3,padding=d,dilation=d) for d in (1,2,4)]);self.dec=nn.Sequential(nn.Conv2d(hidden*3,hidden,3,padding=1),nn.ReLU(),nn.Conv2d(hidden,classes,1));self.model_config={'channels':channels,'classes':classes,'hidden':hidden} def forward(self,x):z=self.enc(x);z=torch.cat([m(z) for m in self.aspp],1);return F.interpolate(self.dec(z),size=x.shape[-2:],mode='bilinear',align_corners=False) def write(p,o):p=Path(p);p.parent.mkdir(parents=True,exist_ok=True);p.write_text(json.dumps(o,indent=2)+'\n')