"""Compact physical-state coupled atmosphere-ocean emulator.""" import json from pathlib import Path import numpy as np import torch from torch import nn import yaml ATM3D=("T","qT","U","V");ATM2D=("Ts","ps");DIAG=("RSW","OLR","USWsfc","ULWsfc","DSWsfc","DLWsfc","LHF","SHF","P","dTWP_adv","tau_u","tau_v") OCEAN3D=("thetao","so","uo","vo");OCEAN2D=("SIC","HI","SST","ZOS") def load_config(root):return yaml.safe_load((Path(root)/"conf/config.yaml").read_text()) def synthetic_state(origin,size,i,step=0): y,x=torch.meshgrid(torch.arange(origin[0],origin[0]+size),torch.arange(origin[1],origin[1]+size),indexing="ij");lat=torch.deg2rad(90-(y+.5));lon=torch.deg2rad(-180+(x+.5));a=torch.arange(46)[:,None,None];o=torch.arange(80)[:,None,None];atm=torch.sin((a%6+1)*lat+.03*a+.05*(i+step))*torch.cos(lon);oce=torch.cos((o%5+1)*lat-.01*o+.02*(i+step))*torch.sin(lon);return atm.float(),oce.float() class ResidualEmulator(nn.Module): def __init__(self,cin,cout,h):super().__init__();self.net=nn.Sequential(nn.Conv2d(cin,h,1),nn.GELU(),nn.Conv2d(h,h,3,padding=1),nn.GELU(),nn.Conv2d(h,cout,1)) def forward(self,x):return self.net(x) class SamudrACE(nn.Module): def __init__(self,hidden_dim=16):super().__init__();self.atmosphere=ResidualEmulator(48,46,hidden_dim);self.ocean=ResidualEmulator(92,80,hidden_dim);self.model_config={"hidden_dim":hidden_dim} def atmosphere_step(self,atm,ocean):return atm+self.atmosphere(torch.cat((atm,ocean[:,[76,0]]),1)) def ocean_step(self,ocean,flux):return ocean+self.ocean(torch.cat((ocean,flux),1)) def forward(self,atm,ocean,n=20):return self.coupled_step(atm,ocean,n) def coupled_step(self,atm,ocean,n=20): flux=[] for _ in range(n):atm=self.atmosphere_step(atm,ocean);flux.append(atm[:,-12:]) ocean=self.ocean_step(ocean,torch.stack(flux).mean(0));return atm,ocean 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")