zhangrenchao's picture
Publish PrecipExtremes-GAN reproduction
ffdf763 verified
Raw
History Blame Contribute Delete
3 kB
from pathlib import Path
import sys,os
import numpy as np,torch
import torch.nn.functional as F
from torch.nn.parallel import DistributedDataParallel
ROOT=Path(__file__).resolve().parents[1];sys.path.insert(0,str(ROOT))
from model.precip_extremes_gan import *
c=load_config(ROOT);rank=int(os.environ.get("RANK",0));world=int(os.environ.get("WORLD_SIZE",1));local=int(os.environ.get("LOCAL_RANK",0));distributed=world>1;use_cuda=torch.cuda.is_available() and torch.cuda.device_count()>=world and c["runtime"]["device"]!="cpu"
if distributed:torch.distributed.init_process_group(backend="nccl" if use_cuda else "gloo")
device=torch.device(f"cuda:{local}" if use_cuda else "cpu");torch.manual_seed(c["seed"]+rank);torch.set_num_threads(2);d=np.load(ROOT/c["data"]["path"]);m=d["split"]=="train";x=torch.from_numpy(d["predictors"][m]).to(device);y=torch.log(torch.from_numpy(d["precipitation"][m])+c["data"]["log_epsilon"]).to(device);net=PrecipExtremesGAN(c["data"]["input_channels"],c["model"]["filters"],c["model"]["noise_channels"]).to(device);critic=PatchCritic(c["data"]["input_channels"]).to(device)
baseline=DistributedDataParallel(net.baseline,device_ids=[local] if use_cuda else None) if distributed else net.baseline;generator=DistributedDataParallel(net.generator,device_ids=[local] if use_cuda else None) if distributed else net.generator;critic_train=DistributedDataParallel(critic,device_ids=[local] if use_cuda else None) if distributed else critic
ob=torch.optim.Adam(baseline.parameters(),lr=c["train"]["baseline_lr"])
for _ in range(c["train"]["baseline_epochs"]): pred=baseline(x,c["data"]["high_grid"]);loss=F.mse_loss(pred,y);ob.zero_grad();loss.backward();ob.step()
og=torch.optim.Adam(generator.parameters(),lr=c["train"]["gan_lr"]);oc=torch.optim.Adam(critic_train.parameters(),lr=c["train"]["gan_lr"])
for _ in range(c["train"]["gan_epochs"]):
with torch.no_grad():base=baseline(x,c["data"]["high_grid"])
fake=base+generator(torch.cat((x,torch.randn(len(x),1,*x.shape[-2:],device=device)),1),c["data"]["high_grid"]);cl=F.softplus(critic_train(x,fake.detach())).mean()+F.softplus(-critic_train(x,y)).mean();oc.zero_grad();cl.backward();oc.step();adv=F.softplus(-critic_train(x,fake)).mean();intensity=F.mse_loss(fake.mean((2,3)),y.mean((2,3)));gl=F.mse_loss(fake.mean(0),y.mean(0))+c["train"]["adversarial_weight"]*adv+intensity;og.zero_grad();gl.backward();og.step()
vals=torch.tensor([float(loss),float(gl),float(cl)],device=device);
if distributed:torch.distributed.all_reduce(vals);vals/=world
if rank==0:
path=ROOT/c["paths"]["checkpoint"];path.parent.mkdir(parents=True,exist_ok=True);torch.save({"model":net.state_dict(),"critic":critic.state_dict(),"model_config":net.model_config,"format_version":c["data"]["format_version"]},path);write_json(ROOT/c["paths"]["training_metrics"],{"baseline_mse":float(vals[0]),"generator_loss":float(vals[1]),"critic_loss":float(vals[2]),"world_size":world});print(path)
if distributed:torch.distributed.destroy_process_group()