WeatherBench / scripts /train.py
zhangrenchao's picture
Publish WeatherBench reproduction
aa529c9 verified
Raw
History Blame Contribute Delete
1.28 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.weatherbench 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"]);d=np.load(ROOT/c["data"]["path"]);ids=np.where(d["split"]==0)[0];base=WeatherBenchCNN(**c["model"]);m=DDP(base) if distributed else base;opt=torch.optim.Adam(m.parameters(),lr=c["train"]["learning_rate"]);losses=[]
for _ in range(c["train"]["epochs"]):
for i in ids[rank::world]:x=torch.tensor(d["input"][i:i+1]);y=torch.tensor(d["target_6h"][i:i+1]);loss=((m(x)-y)**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)
p=ROOT/c["paths"]["checkpoint"]
if rank==0:p.parent.mkdir(parents=True,exist_ok=True);torch.save({"model":base.state_dict(),"model_config":c["model"]},p);write_json(ROOT/c["paths"]["training_metrics"],{"mse":float(v[0]/v[1]),"world_size":world});print(p)
if distributed:dist.destroy_process_group()