SamudrACE / scripts /train.py
zhangrenchao's picture
Publish SamudrACE reproduction
956cd31 verified
Raw
History Blame Contribute Delete
1.38 kB
from pathlib import Path
import sys,os,numpy as np,torch
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
from model.samudrace import *
c=load_config(ROOT);rank=int(os.getenv("RANK",0));world=int(os.getenv("WORLD_SIZE",1));distributed=world>1
if distributed:dist.init_process_group("gloo")
torch.manual_seed(c["seed"]);torch.set_num_threads(2);d=np.load(ROOT/c["data"]["path"]);base=SamudrACE(**c["model"]);m=DDP(base) if distributed else base;opt=torch.optim.AdamW(m.parameters(),lr=c["train"]["learning_rate"]);losses=[]
for i in range(rank,len(d["atmosphere"]),world):a=torch.tensor(d["atmosphere"][i:i+1]);o=torch.tensor(d["ocean"][i:i+1]);pa,po=m(a,o,20);loss=((pa-torch.tensor(d["atmosphere_target"][i:i+1]))**2).mean()+((po-torch.tensor(d["ocean_target"][i:i+1]))**2).mean();opt.zero_grad();loss.backward();opt.step();losses.append(float(loss))
v=torch.tensor([sum(losses),len(losses)],dtype=torch.float64)
if distributed:dist.all_reduce(v)
path=ROOT/c["paths"]["checkpoint"]
if rank==0:path.parent.mkdir(parents=True,exist_ok=True);torch.save({"model":base.state_dict(),"model_config":c["model"]},path);write_json(ROOT/c["paths"]["training_metrics"],{"coupled_mse":float(v[0]/v[1]),"world_size":world});print(path)
if distributed:dist.destroy_process_group()