Download model/samudrace.py from OneScience-Group/SamudrACE: direct link, hf CLI and curl.
- Browser
- Download file 1.95 kB
-
https://huggingface.co/OneScience-Group/SamudrACE/resolve/main/model/samudrace.py
- Command line
-
hf download hf://OneScience-Group/SamudrACE/model/samudrace.py
-
curl -L -o samudrace.py https://huggingface.co/OneScience-Group/SamudrACE/resolve/main/model/samudrace.py
1.95 kB
| """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") | |